mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-29 02:13:55 +08:00
Compare commits
4
Commits
py-test-mqtt
..
main
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
de45f73efd | ||
|
|
1dc46e30d7 | ||
|
|
18e5c47a06 | ||
|
|
4a2c90ec87 |
@@ -1,161 +0,0 @@
|
||||
# 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.
|
||||
@@ -21,14 +21,14 @@
|
||||
<junit.version>5.10.1</junit.version>
|
||||
<druid.version>1.2.20</druid.version>
|
||||
<mybatisplus.version>3.5.17</mybatisplus.version>
|
||||
<hutool.version>5.8.24</hutool.version>
|
||||
<jsoup.version>1.19.1</jsoup.version>
|
||||
<hutool.version>5.8.46</hutool.version>
|
||||
<jsoup.version>1.22.2</jsoup.version>
|
||||
<knife4j.version>4.6.0</knife4j.version>
|
||||
<springdoc.version>2.8.8</springdoc.version>
|
||||
<commons-lang3.version>3.18.0</commons-lang3.version>
|
||||
<commons-lang3.version>3.20.0</commons-lang3.version>
|
||||
<shiro.version>2.0.2</shiro.version>
|
||||
<captcha.version>1.6.2</captcha.version>
|
||||
<guava.version>33.0.0-jre</guava.version>
|
||||
<guava.version>33.6.0-jre</guava.version>
|
||||
<liquibase-core.version>4.20.0</liquibase-core.version>
|
||||
<aliyun-sms-version>4.1.0</aliyun-sms-version>
|
||||
<okio-version>3.4.0</okio-version>
|
||||
|
||||
@@ -106,32 +106,6 @@ 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地址
|
||||
*/
|
||||
|
||||
@@ -1,52 +0,0 @@
|
||||
package xiaozhi.common.utils;
|
||||
|
||||
import lombok.extern.slf4j.Slf4j;
|
||||
|
||||
import java.security.MessageDigest;
|
||||
import java.security.NoSuchAlgorithmException;
|
||||
|
||||
/**
|
||||
* 哈希加密算法的工具类
|
||||
* @author zjy
|
||||
*/
|
||||
@Slf4j
|
||||
public class HashEncryptionUtil {
|
||||
/**
|
||||
* 使用md5进行加密
|
||||
* @param context 被加密的内容
|
||||
* @return 哈希值
|
||||
*/
|
||||
public static String Md5hexDigest(String context){
|
||||
return hexDigest(context,"MD5");
|
||||
}
|
||||
|
||||
/**
|
||||
* 指定哈希算法进行加密
|
||||
* @param context 被加密的内容
|
||||
* @param algorithm 哈希算法
|
||||
* @return 哈希值
|
||||
*/
|
||||
public static String hexDigest(String context,String algorithm ){
|
||||
// 获取MD5算法实例
|
||||
MessageDigest md = null;
|
||||
try {
|
||||
md = MessageDigest.getInstance(algorithm);
|
||||
} catch (NoSuchAlgorithmException e) {
|
||||
log.error("加密失败的算法:{}",algorithm);
|
||||
throw new RuntimeException("加密失败,"+ algorithm +"哈希算法系统不支持");
|
||||
}
|
||||
// 计算智能体id的MD5值
|
||||
byte[] messageDigest = md.digest(context.getBytes());
|
||||
// 将字节数组转换为十六进制字符串
|
||||
StringBuilder hexString = new StringBuilder();
|
||||
for (byte b : messageDigest) {
|
||||
String hex = Integer.toHexString(0xFF & b);
|
||||
if (hex.length() == 1) {
|
||||
hexString.append('0');
|
||||
}
|
||||
hexString.append(hex);
|
||||
}
|
||||
return hexString.toString();
|
||||
}
|
||||
|
||||
}
|
||||
@@ -1,89 +0,0 @@
|
||||
package xiaozhi.common.utils;
|
||||
|
||||
import cn.hutool.core.util.ReUtil;
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
|
||||
import java.time.LocalDateTime;
|
||||
import java.time.ZoneId;
|
||||
import java.util.Date;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.Set;
|
||||
|
||||
/**
|
||||
* 通用工具类
|
||||
*/
|
||||
public class ToolUtil {
|
||||
private static final Logger logger = LoggerFactory.getLogger(ToolUtil.class);
|
||||
|
||||
/**
|
||||
* 对象是否不为空(新增)
|
||||
*/
|
||||
public static boolean isNotEmpty(Object o) {
|
||||
return !isEmpty(o);
|
||||
}
|
||||
|
||||
/**
|
||||
* 对象是否为空
|
||||
*/
|
||||
public static boolean isEmpty(Object o) {
|
||||
if (o == null) {
|
||||
return true;
|
||||
}
|
||||
if (o instanceof String) {
|
||||
if (o.toString().trim().equals("")) {
|
||||
return true;
|
||||
}
|
||||
} else if (o instanceof List) {
|
||||
if (((List) o).size() == 0) {
|
||||
return true;
|
||||
}
|
||||
} else if (o instanceof Map) {
|
||||
if (((Map) o).size() == 0) {
|
||||
return true;
|
||||
}
|
||||
} else if (o instanceof Set) {
|
||||
if (((Set) o).size() == 0) {
|
||||
return true;
|
||||
}
|
||||
} else if (o instanceof Object[]) {
|
||||
if (((Object[]) o).length == 0) {
|
||||
return true;
|
||||
}
|
||||
} else if (o instanceof int[]) {
|
||||
if (((int[]) o).length == 0) {
|
||||
return true;
|
||||
}
|
||||
} else if (o instanceof long[]) {
|
||||
if (((long[]) o).length == 0) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
/**
|
||||
* 对象组中是否存在空对象
|
||||
*/
|
||||
public static boolean isOneEmpty(Object... os) {
|
||||
for (Object o : os) {
|
||||
if (isEmpty(o)) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
/**
|
||||
* 对象组中是否全是空对象
|
||||
*/
|
||||
public static boolean isAllEmpty(Object... os) {
|
||||
for (Object o : os) {
|
||||
if (!isEmpty(o)) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
return true;
|
||||
}
|
||||
}
|
||||
+2
-2
@@ -5,6 +5,7 @@ import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.stream.Collectors;
|
||||
|
||||
import cn.hutool.core.collection.CollUtil;
|
||||
import cn.hutool.core.collection.ListUtil;
|
||||
import lombok.RequiredArgsConstructor;
|
||||
import org.springframework.stereotype.Service;
|
||||
@@ -20,7 +21,6 @@ import xiaozhi.common.constant.Constant;
|
||||
import xiaozhi.common.page.PageData;
|
||||
import xiaozhi.common.utils.ConvertUtils;
|
||||
import xiaozhi.common.utils.JsonUtils;
|
||||
import xiaozhi.common.utils.ToolUtil;
|
||||
import xiaozhi.modules.agent.Enums.AgentChatHistoryType;
|
||||
import xiaozhi.modules.agent.dao.AiAgentChatHistoryDao;
|
||||
import xiaozhi.modules.agent.dto.AgentChatHistoryDTO;
|
||||
@@ -107,7 +107,7 @@ public class AgentChatHistoryServiceImpl extends CrudRepository<AiAgentChatHisto
|
||||
if (deleteAudio) {
|
||||
// 分批删除音频,避免超时
|
||||
List<String> audioIds = baseMapper.getAudioIdsByAgentId(agentId);
|
||||
if (ToolUtil.isNotEmpty(audioIds)) {
|
||||
if (CollUtil.isNotEmpty(audioIds)) {
|
||||
// 每批删除1000条
|
||||
List<List<String>> batch = ListUtil.split(audioIds, 1000);
|
||||
batch.forEach(dataList -> {
|
||||
|
||||
+2
-2
@@ -12,11 +12,11 @@ import java.util.stream.Collectors;
|
||||
import org.apache.commons.lang3.StringUtils;
|
||||
import org.springframework.stereotype.Service;
|
||||
|
||||
import cn.hutool.crypto.digest.DigestUtil;
|
||||
import lombok.AllArgsConstructor;
|
||||
import lombok.extern.slf4j.Slf4j;
|
||||
import xiaozhi.common.constant.Constant;
|
||||
import xiaozhi.common.utils.AESUtils;
|
||||
import xiaozhi.common.utils.HashEncryptionUtil;
|
||||
import xiaozhi.common.utils.JsonUtils;
|
||||
import xiaozhi.modules.agent.Enums.XiaoZhiMcpJsonRpcJson;
|
||||
import xiaozhi.modules.agent.service.AgentMcpAccessPointService;
|
||||
@@ -226,7 +226,7 @@ public class AgentMcpAccessPointServiceImpl implements AgentMcpAccessPointServic
|
||||
*/
|
||||
private static String encryptToken(String agentId, String key) {
|
||||
// 使用md5对智能体id进行加密
|
||||
String md5 = HashEncryptionUtil.Md5hexDigest(agentId);
|
||||
String md5 = DigestUtil.md5Hex(agentId);
|
||||
// aes需要加密文本
|
||||
String json = "{\"agentId\": \"%s\"}".formatted(md5);
|
||||
// 加密后成token值
|
||||
|
||||
+4
-4
@@ -18,6 +18,7 @@ import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper;
|
||||
import com.baomidou.mybatisplus.core.metadata.IPage;
|
||||
import com.baomidou.mybatisplus.extension.repository.IRepository;
|
||||
|
||||
import cn.hutool.core.collection.CollUtil;
|
||||
import lombok.AllArgsConstructor;
|
||||
import xiaozhi.common.constant.Constant;
|
||||
import xiaozhi.common.exception.ErrorCode;
|
||||
@@ -29,7 +30,6 @@ import xiaozhi.common.service.impl.BaseServiceImpl;
|
||||
import xiaozhi.common.user.UserDetail;
|
||||
import xiaozhi.common.utils.ConvertUtils;
|
||||
import xiaozhi.common.utils.JsonUtils;
|
||||
import xiaozhi.common.utils.ToolUtil;
|
||||
import xiaozhi.modules.agent.dao.AgentDao;
|
||||
import xiaozhi.modules.agent.dao.AgentTagDao;
|
||||
import xiaozhi.modules.agent.dto.AgentCreateDTO;
|
||||
@@ -243,13 +243,13 @@ public class AgentServiceImpl extends BaseServiceImpl<AgentDao, AgentEntity> imp
|
||||
.map(DeviceEntity::getAgentId)
|
||||
.distinct()
|
||||
.collect(Collectors.toList());
|
||||
if (ToolUtil.isNotEmpty(agentIds)) {
|
||||
if (CollUtil.isNotEmpty(agentIds)) {
|
||||
w.or().in("id", agentIds);
|
||||
}
|
||||
|
||||
// 按标签名搜索
|
||||
List<String> tagAgentIds = agentTagService.getAgentIdsByTagName(keyword);
|
||||
if (ToolUtil.isNotEmpty(tagAgentIds)) {
|
||||
if (CollUtil.isNotEmpty(tagAgentIds)) {
|
||||
w.or().in("id", tagAgentIds);
|
||||
}
|
||||
});
|
||||
@@ -291,7 +291,7 @@ public class AgentServiceImpl extends BaseServiceImpl<AgentDao, AgentEntity> imp
|
||||
|
||||
// 获取标签列表
|
||||
List<AgentTagEntity> tags = agentTagDao.selectByAgentId(agent.getId());
|
||||
if (ToolUtil.isNotEmpty(tags)) {
|
||||
if (CollUtil.isNotEmpty(tags)) {
|
||||
dto.setTags(tags.stream().map(this::convertTagToDTO).collect(Collectors.toList()));
|
||||
}
|
||||
|
||||
|
||||
+2
-46
@@ -1,7 +1,6 @@
|
||||
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;
|
||||
@@ -79,8 +78,8 @@ public class OTAController {
|
||||
@Hidden
|
||||
public ResponseEntity<String> getOTA() {
|
||||
String mqttUdpConfig = sysParamsService.getValue(Constant.SERVER_MQTT_GATEWAY, false);
|
||||
if (!isNativeMqttReady() && !isConfiguredValue(mqttUdpConfig)) {
|
||||
return ResponseEntity.ok("OTA接口不正常,缺少mqtt_gateway地址或未启用原生MQTT,请登录智控台,在参数管理找到【server.mqtt_gateway】或【mqtt_server.enabled/protocols.mqtt_enabled】配置");
|
||||
if (StringUtils.isBlank(mqttUdpConfig)) {
|
||||
return ResponseEntity.ok("OTA接口不正常,缺少mqtt_gateway地址,请登录智控台,在参数管理找到【server.mqtt_gateway】配置");
|
||||
}
|
||||
String wsUrl = sysParamsService.getValue(Constant.SERVER_WEBSOCKET, true);
|
||||
if (StringUtils.isBlank(wsUrl) || wsUrl.equals("null")) {
|
||||
@@ -93,49 +92,6 @@ 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();
|
||||
|
||||
+4
-11
@@ -5,8 +5,6 @@ import java.io.IOException;
|
||||
import java.nio.file.Files;
|
||||
import java.nio.file.Path;
|
||||
import java.nio.file.Paths;
|
||||
import java.security.MessageDigest;
|
||||
import java.security.NoSuchAlgorithmException;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.Optional;
|
||||
@@ -31,6 +29,7 @@ import org.springframework.web.bind.annotation.RequestParam;
|
||||
import org.springframework.web.bind.annotation.RestController;
|
||||
import org.springframework.web.multipart.MultipartFile;
|
||||
|
||||
import cn.hutool.crypto.digest.DigestUtil;
|
||||
import io.swagger.v3.oas.annotations.Operation;
|
||||
import io.swagger.v3.oas.annotations.Parameter;
|
||||
import io.swagger.v3.oas.annotations.Parameters;
|
||||
@@ -288,7 +287,7 @@ public class OTAMagController {
|
||||
|
||||
// 返回文件路径
|
||||
return new Result<String>().ok(filePath.toString());
|
||||
} catch (IOException | NoSuchAlgorithmException e) {
|
||||
} catch (IOException e) {
|
||||
return new Result<String>().error("文件上传失败:" + e.getMessage());
|
||||
}
|
||||
}
|
||||
@@ -329,13 +328,7 @@ public class OTAMagController {
|
||||
return result;
|
||||
}
|
||||
|
||||
private String calculateMD5(MultipartFile file) throws IOException, NoSuchAlgorithmException {
|
||||
MessageDigest md = MessageDigest.getInstance("MD5");
|
||||
byte[] digest = md.digest(file.getBytes());
|
||||
StringBuilder sb = new StringBuilder();
|
||||
for (byte b : digest) {
|
||||
sb.append(String.format("%02x", b));
|
||||
}
|
||||
return sb.toString();
|
||||
private String calculateMD5(MultipartFile file) throws IOException {
|
||||
return DigestUtil.md5Hex(file.getBytes());
|
||||
}
|
||||
}
|
||||
|
||||
+25
-74
@@ -18,7 +18,6 @@ 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;
|
||||
@@ -60,14 +59,7 @@ public class DeviceAddressBookServiceImpl implements DeviceAddressBookService {
|
||||
Map<String, Map<String, String>> allBooks = getAllAddressBooks();
|
||||
|
||||
if (isAnswer) {
|
||||
DeviceEntity callerDevice =
|
||||
deviceService.getDeviceByMacAddress(callerMac);
|
||||
if (callerDevice == null) {
|
||||
return errorResult("接听失败,设备信息不存在");
|
||||
}
|
||||
return postCallAccept(
|
||||
buildMqttClientId(callerDevice),
|
||||
Map.of("mac", callerMac));
|
||||
return postToMqtt("/api/call/accept", Map.of("mac", callerMac), "接听");
|
||||
}
|
||||
|
||||
// 主动呼叫模式
|
||||
@@ -87,14 +79,6 @@ 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;
|
||||
@@ -102,19 +86,15 @@ public class DeviceAddressBookServiceImpl implements DeviceAddressBookService {
|
||||
callerNickname = targetBook.get(callerMac.toLowerCase());
|
||||
}
|
||||
if (StringUtils.isBlank(callerNickname)) {
|
||||
callerNickname = callerDevice.getAlias();
|
||||
callerNickname = deviceService.getDeviceByMacAddress(callerMac).getAlias();
|
||||
if (StringUtils.isBlank(callerNickname)) {
|
||||
callerNickname = formatMacAsDeviceName(callerMac);
|
||||
}
|
||||
}
|
||||
|
||||
return postCallRequest(
|
||||
buildMqttClientId(callerDevice),
|
||||
buildMqttClientId(targetDevice),
|
||||
Map.of(
|
||||
"caller_mac", callerMac,
|
||||
"target_mac", targetMac,
|
||||
"caller_nickname", callerNickname));
|
||||
return postToMqtt("/api/call/request",
|
||||
Map.of("caller_mac", callerMac, "target_mac", targetMac, "caller_nickname", callerNickname),
|
||||
"呼叫");
|
||||
}
|
||||
|
||||
@Override
|
||||
@@ -215,50 +195,32 @@ public class DeviceAddressBookServiceImpl implements DeviceAddressBookService {
|
||||
return result;
|
||||
}
|
||||
|
||||
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) {
|
||||
private Map<String, Object> postToMqtt(String path, Map<String, Object> body, String action) {
|
||||
Map<String, Object> result = new HashMap<>();
|
||||
result.put("status", "error");
|
||||
if (response == null || StringUtils.isBlank(response.body())) {
|
||||
result.put("message", action + "失败,MQTT管理配置缺失");
|
||||
|
||||
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 + "失败,网关配置缺失");
|
||||
return result;
|
||||
}
|
||||
|
||||
try {
|
||||
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"));
|
||||
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"));
|
||||
}
|
||||
return result;
|
||||
} catch (Exception e) {
|
||||
@@ -267,17 +229,6 @@ 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;
|
||||
|
||||
+105
-316
@@ -5,13 +5,11 @@ 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;
|
||||
@@ -33,6 +31,7 @@ import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper;
|
||||
import com.baomidou.mybatisplus.core.conditions.update.UpdateWrapper;
|
||||
import com.baomidou.mybatisplus.core.metadata.IPage;
|
||||
|
||||
import cn.hutool.core.collection.CollUtil;
|
||||
import cn.hutool.core.map.MapUtil;
|
||||
import cn.hutool.core.util.RandomUtil;
|
||||
import cn.hutool.core.util.StrUtil;
|
||||
@@ -53,7 +52,6 @@ import xiaozhi.common.user.UserDetail;
|
||||
import xiaozhi.common.utils.ConvertUtils;
|
||||
import xiaozhi.common.utils.DateUtils;
|
||||
import xiaozhi.common.utils.JsonUtils;
|
||||
import xiaozhi.common.utils.ToolUtil;
|
||||
import xiaozhi.modules.device.dao.DeviceDao;
|
||||
import xiaozhi.modules.device.dto.DeviceManualAddDTO;
|
||||
import xiaozhi.modules.device.dto.DevicePageUserDTO;
|
||||
@@ -105,15 +103,15 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
|
||||
throw new RenException(ErrorCode.ACTIVATION_CODE_EMPTY);
|
||||
}
|
||||
String deviceKey = RedisKeys.getOtaActivationCode(activationCode);
|
||||
Object cacheDeviceId = redisUtils.get(deviceKey);
|
||||
if (ToolUtil.isEmpty(cacheDeviceId)) {
|
||||
String cacheDeviceId = (String) redisUtils.get(deviceKey);
|
||||
if (StringUtils.isBlank(cacheDeviceId)) {
|
||||
throw new RenException(ErrorCode.ACTIVATION_CODE_ERROR);
|
||||
}
|
||||
String deviceId = (String) cacheDeviceId;
|
||||
String deviceId = cacheDeviceId;
|
||||
String safeDeviceId = deviceId.replace(":", "_").toLowerCase();
|
||||
String cacheDeviceKey = RedisKeys.getOtaDeviceActivationInfo(safeDeviceId);
|
||||
Map<String, Object> cacheMap = JsonUtils.toStringObjectMap(redisUtils.get(cacheDeviceKey));
|
||||
if (ToolUtil.isEmpty(cacheMap)) {
|
||||
if (MapUtil.isEmpty(cacheMap)) {
|
||||
throw new RenException(ErrorCode.ACTIVATION_CODE_ERROR);
|
||||
}
|
||||
String cachedCode = (String) cacheMap.get("activation_code");
|
||||
@@ -159,20 +157,32 @@ 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(device -> MqttClientId.build(
|
||||
device.getBoard(), device.getMacAddress()))
|
||||
.collect(Collectors.toSet());
|
||||
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());
|
||||
|
||||
// 构建请求入参
|
||||
if (ToolUtil.isNotEmpty(deviceIds)) {
|
||||
return createMqttManagementRouter()
|
||||
.getMergedStatus(deviceIds);
|
||||
Map<String, Set<String>> params = MapUtil
|
||||
.builder(new HashMap<String, Set<String>>())
|
||||
.put("clientIds", deviceIds).build();
|
||||
|
||||
if (CollUtil.isNotEmpty(deviceIds)) {
|
||||
return postToMqttGateway(url, params);
|
||||
}
|
||||
// 返回响应
|
||||
return "";
|
||||
@@ -184,15 +194,6 @@ 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) {
|
||||
@@ -203,7 +204,8 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
|
||||
} else {
|
||||
// 只有在设备已绑定且明确开启自动升级时才返回固件升级信息
|
||||
if (Integer.valueOf(1).equals(deviceById.getAutoUpdate())) {
|
||||
DeviceReportRespDTO.Firmware firmware = buildFirmwareInfo(reportedBoard,
|
||||
String type = deviceReport.getBoard() == null ? null : deviceReport.getBoard().getType();
|
||||
DeviceReportRespDTO.Firmware firmware = buildFirmwareInfo(type,
|
||||
deviceReport.getApplication() == null ? null : deviceReport.getApplication().getVersion());
|
||||
response.setFirmware(firmware);
|
||||
}
|
||||
@@ -246,43 +248,20 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
|
||||
|
||||
response.setWebsocket(websocket);
|
||||
|
||||
// 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()) {
|
||||
// 添加MQTT UDP配置
|
||||
// 从系统参数获取MQTT Gateway地址,仅在配置有效时使用
|
||||
String mqttUdpConfig = sysParamsService.getValue(Constant.SERVER_MQTT_GATEWAY, true);
|
||||
if (mqttUdpConfig != null && !mqttUdpConfig.equals("null") && !mqttUdpConfig.isEmpty()) {
|
||||
try {
|
||||
DeviceReportRespDTO.MQTT mqtt = buildNativeMqttConfig(macAddress, groupId, clientId);
|
||||
String endpoint = buildNativeMqttEndpoint();
|
||||
if (mqtt != null && StringUtils.isNotBlank(endpoint)) {
|
||||
mqtt.setEndpoint(endpoint);
|
||||
String groupId = deviceById != null && deviceById.getBoard() != null ? deviceById.getBoard()
|
||||
: "GID_default";
|
||||
DeviceReportRespDTO.MQTT mqtt = buildMqttConfig(macAddress, groupId);
|
||||
if (mqtt != null) {
|
||||
mqtt.setEndpoint(mqttUdpConfig);
|
||||
response.setMqtt(mqtt);
|
||||
mqttConfigured = true;
|
||||
}
|
||||
} catch (Exception e) {
|
||||
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());
|
||||
}
|
||||
log.error("生成MQTT配置失败: {}", e.getMessage());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -654,210 +633,6 @@ 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配置信息
|
||||
*
|
||||
@@ -875,10 +650,28 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
|
||||
}
|
||||
|
||||
// 构建客户端ID格式:groupId@@@macAddress@@@uuid
|
||||
String deviceIdSafeStr = MqttClientId.normalizeDeviceId(macAddress);
|
||||
String mqttClientId = MqttClientId.build(groupId, macAddress);
|
||||
String groupIdSafeStr = groupId.replace(":", "_");
|
||||
String deviceIdSafeStr = macAddress.replace(":", "_");
|
||||
String mqttClientId = String.format("%s@@@%s@@@%s", groupIdSafeStr, deviceIdSafeStr, deviceIdSafeStr);
|
||||
|
||||
String username = buildMqttUsername();
|
||||
// 构建用户数据(包含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 password = generatePasswordSignature(mqttClientId + "|" + username, signatureKey);
|
||||
@@ -894,14 +687,23 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
|
||||
return mqtt;
|
||||
}
|
||||
|
||||
MqttManagementRouter createMqttManagementRouter() {
|
||||
return new MqttManagementRouter(
|
||||
new MqttManagementEndpointResolver(sysParamsService),
|
||||
new MqttManagementHttpClient());
|
||||
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());
|
||||
}
|
||||
|
||||
@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) {
|
||||
@@ -915,20 +717,19 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
|
||||
}
|
||||
|
||||
// 构建clientId
|
||||
String clientId = MqttClientId.build(
|
||||
device.getBoard(), device.getMacAddress());
|
||||
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);
|
||||
|
||||
// 存储所有工具列表
|
||||
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 (pageCount++ < 32) {
|
||||
while (true) {
|
||||
// 构建params
|
||||
Map<String, Object> paramsMap = MapUtil.builder(new HashMap<String, Object>())
|
||||
.put("withUserTools", true)
|
||||
@@ -953,44 +754,26 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
|
||||
.put("payload", payload)
|
||||
.build();
|
||||
|
||||
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();
|
||||
}
|
||||
String resultMessage = postToMqttGateway(url, requestBody);
|
||||
|
||||
// 解析响应
|
||||
if (StringUtils.isBlank(resultMessage)) {
|
||||
return null;
|
||||
break;
|
||||
}
|
||||
|
||||
JSONObject jsonObject;
|
||||
try {
|
||||
jsonObject = JSONUtil.parseObj(resultMessage);
|
||||
} catch (RuntimeException e) {
|
||||
return null;
|
||||
}
|
||||
JSONObject jsonObject = JSONUtil.parseObj(resultMessage);
|
||||
if (!jsonObject.getBool("success", false)) {
|
||||
return null;
|
||||
break;
|
||||
}
|
||||
|
||||
JSONObject data = jsonObject.getJSONObject("data");
|
||||
if (data == null) {
|
||||
return null;
|
||||
break;
|
||||
}
|
||||
|
||||
// 获取当前页的工具列表
|
||||
JSONArray tools = data.getJSONArray("tools");
|
||||
if (tools == null) {
|
||||
return null;
|
||||
}
|
||||
if (!tools.isEmpty()) {
|
||||
if (tools != null && !tools.isEmpty()) {
|
||||
allTools.addAll(tools);
|
||||
}
|
||||
|
||||
@@ -998,23 +781,29 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
|
||||
String nextCursor = data.getStr("nextCursor");
|
||||
if (StringUtils.isBlank(nextCursor)) {
|
||||
// 没有下一页了
|
||||
Map<String, Object> resultData = new HashMap<>();
|
||||
resultData.put("tools", allTools);
|
||||
return resultData;
|
||||
}
|
||||
if (!seenCursors.add(nextCursor)) {
|
||||
log.warn("MQTT设备工具列表返回重复cursor,终止分页: {}",
|
||||
nextCursor);
|
||||
return null;
|
||||
break;
|
||||
}
|
||||
cursor = nextCursor;
|
||||
}
|
||||
|
||||
return null;
|
||||
// 构建返回结果
|
||||
if (allTools.isEmpty()) {
|
||||
return null;
|
||||
}
|
||||
|
||||
Map<String, Object> resultData = new HashMap<>();
|
||||
resultData.put("tools", allTools);
|
||||
return resultData;
|
||||
}
|
||||
|
||||
@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) {
|
||||
@@ -1028,8 +817,12 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
|
||||
}
|
||||
|
||||
// 构建clientId
|
||||
String clientId = MqttClientId.build(
|
||||
device.getBoard(), device.getMacAddress());
|
||||
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);
|
||||
|
||||
// 构建请求体
|
||||
Map<String, Object> params = MapUtil
|
||||
@@ -1052,11 +845,7 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
|
||||
.put("payload", payload)
|
||||
.build();
|
||||
|
||||
MqttManagementHttpClient.Response response =
|
||||
createMqttManagementRouter()
|
||||
.sendMutatingCommand(clientId, requestBody);
|
||||
String resultMessage =
|
||||
response == null ? null : response.body();
|
||||
String resultMessage = postToMqttGateway(url, requestBody);
|
||||
|
||||
// 解析响应
|
||||
if (StringUtils.isNotBlank(resultMessage)) {
|
||||
|
||||
@@ -1,28 +0,0 @@
|
||||
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("-", "_");
|
||||
}
|
||||
}
|
||||
+4
-15
@@ -27,8 +27,10 @@ final class MqttGatewayAuthorization {
|
||||
}
|
||||
|
||||
static String postJson(String url, String jsonBody, String signatureKey, Instant now, int timeoutMillis) {
|
||||
GatewayResponse response = postJsonResponse(
|
||||
url, jsonBody, signatureKey, now, timeoutMillis);
|
||||
GatewayResponse response = executeWithDateFallback(
|
||||
signatureKey,
|
||||
now,
|
||||
token -> executeRequest(url, jsonBody, token, timeoutMillis));
|
||||
|
||||
if (response.statusCode() < 200 || response.statusCode() >= 300) {
|
||||
throw new GatewayRequestException(
|
||||
@@ -38,19 +40,6 @@ 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
@@ -1,183 +0,0 @@
|
||||
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
@@ -1,56 +0,0 @@
|
||||
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
@@ -1,431 +0,0 @@
|
||||
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) {
|
||||
}
|
||||
}
|
||||
+2
-2
@@ -15,6 +15,7 @@ import org.springframework.web.bind.annotation.RequestMapping;
|
||||
import org.springframework.web.bind.annotation.RequestParam;
|
||||
import org.springframework.web.bind.annotation.RestController;
|
||||
|
||||
import cn.hutool.core.collection.CollUtil;
|
||||
import io.swagger.v3.oas.annotations.Operation;
|
||||
import io.swagger.v3.oas.annotations.Parameter;
|
||||
import io.swagger.v3.oas.annotations.tags.Tag;
|
||||
@@ -23,7 +24,6 @@ import xiaozhi.common.exception.ErrorCode;
|
||||
import xiaozhi.common.exception.RenException;
|
||||
import xiaozhi.common.page.PageData;
|
||||
import xiaozhi.common.utils.Result;
|
||||
import xiaozhi.common.utils.ToolUtil;
|
||||
import xiaozhi.modules.knowledge.dto.KnowledgeBaseDTO;
|
||||
import xiaozhi.modules.knowledge.service.KnowledgeBaseService;
|
||||
import xiaozhi.modules.knowledge.service.KnowledgeManagerService;
|
||||
@@ -140,7 +140,7 @@ public class KnowledgeBaseController {
|
||||
List<String> idList = Arrays.asList(ids.split(","));
|
||||
List<KnowledgeBaseDTO> knowledgeBaseDTOs = Optional.ofNullable(knowledgeBaseService.getByDatasetIdList(idList))
|
||||
.orElseGet(ArrayList::new);
|
||||
if (ToolUtil.isNotEmpty(knowledgeBaseDTOs)) {
|
||||
if (CollUtil.isNotEmpty(knowledgeBaseDTOs)) {
|
||||
knowledgeBaseDTOs.forEach(item -> {
|
||||
// 检查权限:用户只能删除自己创建的知识库
|
||||
if (item.getCreator() == null || !item.getCreator().equals(currentUserId)) {
|
||||
|
||||
+4
-20
@@ -1,8 +1,7 @@
|
||||
package xiaozhi.modules.security.service.impl;
|
||||
|
||||
import java.io.IOException;
|
||||
import java.util.Random;
|
||||
import java.util.concurrent.TimeUnit;
|
||||
import java.time.Duration;
|
||||
|
||||
import org.apache.commons.lang3.StringUtils;
|
||||
import org.springframework.beans.factory.annotation.Value;
|
||||
@@ -13,6 +12,7 @@ import com.google.common.cache.CacheBuilder;
|
||||
import com.wf.captcha.SpecCaptcha;
|
||||
import com.wf.captcha.base.Captcha;
|
||||
|
||||
import cn.hutool.core.util.RandomUtil;
|
||||
import jakarta.annotation.Resource;
|
||||
import jakarta.servlet.http.HttpServletResponse;
|
||||
import xiaozhi.common.constant.Constant;
|
||||
@@ -41,7 +41,7 @@ public class CaptchaServiceImpl implements CaptchaService {
|
||||
* Local Cache 5分钟过期
|
||||
*/
|
||||
Cache<String, String> localCache = CacheBuilder.newBuilder().maximumSize(1000)
|
||||
.expireAfterAccess(5, TimeUnit.MINUTES).build();
|
||||
.expireAfterAccess(Duration.ofMinutes(5)).build();
|
||||
|
||||
@Override
|
||||
public void create(HttpServletResponse response, String uuid) throws IOException {
|
||||
@@ -113,7 +113,7 @@ public class CaptchaServiceImpl implements CaptchaService {
|
||||
}
|
||||
|
||||
String key = RedisKeys.getSMSValidateCodeKey(phone);
|
||||
String validateCodes = generateValidateCode(6);
|
||||
String validateCodes = RandomUtil.randomNumbers(6);
|
||||
|
||||
// 设置验证码
|
||||
setCache(key, validateCodes);
|
||||
@@ -135,22 +135,6 @@ public class CaptchaServiceImpl implements CaptchaService {
|
||||
return validate(key, code, delete);
|
||||
}
|
||||
|
||||
/**
|
||||
* 生成指定数量的随机数验证码
|
||||
*
|
||||
* @param length 数量
|
||||
* @return 随机码
|
||||
*/
|
||||
private String generateValidateCode(Integer length) {
|
||||
String chars = "0123456789"; // 字符范围可以自定义:数字
|
||||
Random random = new Random();
|
||||
StringBuilder code = new StringBuilder();
|
||||
for (int i = 0; i < length; i++) {
|
||||
code.append(chars.charAt(random.nextInt(chars.length())));
|
||||
}
|
||||
return code.toString();
|
||||
}
|
||||
|
||||
private void setCache(String key, String value) {
|
||||
if (open) {
|
||||
key = RedisKeys.getCaptchaKey(key);
|
||||
|
||||
+1
-5
@@ -85,7 +85,6 @@ 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);
|
||||
@@ -278,10 +277,7 @@ public class SysParamsController {
|
||||
|
||||
// 校验mqtt密钥长度和复杂度
|
||||
private void validateMqttSecretLength(String paramCode, String secret) {
|
||||
if (!paramCode.equals(Constant.SERVER_MQTT_SECRET)
|
||||
&& !paramCode.equals(Constant.MQTT_SERVER_SIGNATURE_KEY)
|
||||
&& !paramCode.equals(
|
||||
Constant.MQTT_SERVER_MANAGER_API_SECRET)) {
|
||||
if (!paramCode.equals(Constant.SERVER_MQTT_SECRET)) {
|
||||
return;
|
||||
}
|
||||
if (StringUtils.isBlank(secret) || secret.equals("null")) {
|
||||
|
||||
+3
-3
@@ -12,6 +12,7 @@ import com.baomidou.mybatisplus.core.conditions.query.LambdaQueryWrapper;
|
||||
import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper;
|
||||
import com.baomidou.mybatisplus.core.metadata.IPage;
|
||||
|
||||
import cn.hutool.core.collection.CollUtil;
|
||||
import lombok.AllArgsConstructor;
|
||||
import xiaozhi.common.exception.RenException;
|
||||
import xiaozhi.common.exception.ErrorCode;
|
||||
@@ -21,7 +22,6 @@ import xiaozhi.common.redis.RedisUtils;
|
||||
import xiaozhi.common.service.impl.BaseServiceImpl;
|
||||
import xiaozhi.common.utils.ConvertUtils;
|
||||
import xiaozhi.common.utils.JsonUtils;
|
||||
import xiaozhi.common.utils.ToolUtil;
|
||||
import xiaozhi.modules.sys.dao.SysDictDataDao;
|
||||
import xiaozhi.modules.sys.dao.SysUserDao;
|
||||
import xiaozhi.modules.sys.dto.SysDictDataDTO;
|
||||
@@ -104,13 +104,13 @@ public class SysDictDataServiceImpl extends BaseServiceImpl<SysDictDataDao, SysD
|
||||
@Transactional(rollbackFor = Exception.class)
|
||||
public void delete(Long[] ids) {
|
||||
List<Long> idList = Arrays.asList(ids);
|
||||
if (ToolUtil.isNotEmpty(idList)) {
|
||||
if (CollUtil.isNotEmpty(idList)) {
|
||||
//批量删除redis字典
|
||||
List<String> redisKeyList = new ArrayList<>();
|
||||
//批量获取字典类型
|
||||
List<String> dictTypeList = Optional.ofNullable(baseDao.getDictTypesByIdList(idList)).orElseGet(ArrayList::new);
|
||||
dictTypeList.forEach(dictType -> redisKeyList.add(RedisKeys.getDictDataByTypeKey(dictType)));
|
||||
if (ToolUtil.isNotEmpty(redisKeyList)) {
|
||||
if (CollUtil.isNotEmpty(redisKeyList)) {
|
||||
//清除缓存
|
||||
redisUtils.delete(redisKeyList);
|
||||
}
|
||||
|
||||
+1
-6
@@ -205,14 +205,9 @@ public class SysParamsServiceImpl extends BaseServiceImpl<SysParamsDao, SysParam
|
||||
public void initServerSecret() {
|
||||
// 获取服务器密钥
|
||||
String secretParam = getValue(Constant.SERVER_SECRET, false);
|
||||
if (StringUtils.isBlank(secretParam)
|
||||
|| "null".equalsIgnoreCase(secretParam.trim())) {
|
||||
if (StringUtils.isBlank(secretParam) || "null".equals(secretParam)) {
|
||||
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密钥对
|
||||
|
||||
@@ -34,7 +34,7 @@ public class TimbreDataDTO {
|
||||
|
||||
@Schema(description = "排序")
|
||||
@Min(value = 0, message = "{sort.number}")
|
||||
private long sort;
|
||||
private Long sort;
|
||||
|
||||
@Schema(description = "对应 TTS 模型主键")
|
||||
@NotBlank(message = "{timbre.ttsModelId.require}")
|
||||
|
||||
@@ -3,6 +3,7 @@ package xiaozhi.modules.timbre.entity;
|
||||
import java.util.Date;
|
||||
|
||||
import com.baomidou.mybatisplus.annotation.FieldFill;
|
||||
import com.baomidou.mybatisplus.annotation.FieldStrategy;
|
||||
import com.baomidou.mybatisplus.annotation.TableField;
|
||||
import com.baomidou.mybatisplus.annotation.TableName;
|
||||
|
||||
@@ -41,7 +42,8 @@ public class TimbreEntity {
|
||||
private String referenceText;
|
||||
|
||||
@Schema(description = "排序")
|
||||
private long sort;
|
||||
@TableField(updateStrategy = FieldStrategy.NOT_NULL)
|
||||
private Long sort;
|
||||
|
||||
@Schema(description = "对应 TTS 模型主键")
|
||||
private String ttsModelId;
|
||||
|
||||
+3
@@ -99,6 +99,9 @@ public class TimbreServiceImpl extends BaseServiceImpl<TimbreDao, TimbreEntity>
|
||||
@Transactional(rollbackFor = Exception.class)
|
||||
public void save(TimbreDataDTO dto) {
|
||||
isTtsModelId(dto.getTtsModelId());
|
||||
if (dto.getSort() == null) {
|
||||
dto.setSort(0L);
|
||||
}
|
||||
TimbreEntity timbreEntity = ConvertUtils.sourceToTarget(dto, TimbreEntity.class);
|
||||
baseDao.insert(timbreEntity);
|
||||
}
|
||||
|
||||
@@ -32,7 +32,7 @@ public class TimbreDetailsVO implements Serializable {
|
||||
private String referenceText;
|
||||
|
||||
@Schema(description = "排序")
|
||||
private long sort;
|
||||
private Long sort;
|
||||
|
||||
@Schema(description = "对应 TTS 模型主键")
|
||||
private String ttsModelId;
|
||||
|
||||
+2
-2
@@ -17,6 +17,7 @@ import com.baomidou.mybatisplus.core.metadata.IPage;
|
||||
import com.fasterxml.jackson.core.type.TypeReference;
|
||||
import com.fasterxml.jackson.databind.ObjectMapper;
|
||||
|
||||
import cn.hutool.core.collection.CollUtil;
|
||||
import lombok.RequiredArgsConstructor;
|
||||
import lombok.extern.slf4j.Slf4j;
|
||||
import xiaozhi.common.constant.Constant;
|
||||
@@ -26,7 +27,6 @@ import xiaozhi.common.page.PageData;
|
||||
import xiaozhi.common.service.impl.BaseServiceImpl;
|
||||
import xiaozhi.common.utils.ConvertUtils;
|
||||
import xiaozhi.common.utils.DateUtils;
|
||||
import xiaozhi.common.utils.ToolUtil;
|
||||
import xiaozhi.modules.model.entity.ModelConfigEntity;
|
||||
import xiaozhi.modules.model.service.ModelConfigService;
|
||||
import xiaozhi.modules.sys.dao.SysUserDao;
|
||||
@@ -120,7 +120,7 @@ public class VoiceCloneServiceImpl extends BaseServiceImpl<VoiceCloneDao, VoiceC
|
||||
entity.setTrainStatus(0); // 默认训练中
|
||||
batchInsertList.add(entity);
|
||||
}
|
||||
if (ToolUtil.isNotEmpty(batchInsertList)) {
|
||||
if (CollUtil.isNotEmpty(batchInsertList)) {
|
||||
insertBatch(batchInsertList);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,15 +0,0 @@
|
||||
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签名密钥');
|
||||
@@ -1,9 +0,0 @@
|
||||
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连接清理超时(秒)');
|
||||
@@ -1,11 +0,0 @@
|
||||
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,26 +711,3 @@ 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
|
||||
|
||||
+74
@@ -2,16 +2,20 @@ package xiaozhi.modules.timbre.service.impl;
|
||||
|
||||
import static org.junit.jupiter.api.Assertions.assertEquals;
|
||||
import static org.junit.jupiter.api.Assertions.assertNull;
|
||||
import static org.mockito.ArgumentMatchers.argThat;
|
||||
import static org.mockito.Mockito.mock;
|
||||
import static org.mockito.Mockito.never;
|
||||
import static org.mockito.Mockito.verify;
|
||||
import static org.mockito.Mockito.when;
|
||||
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.test.util.ReflectionTestUtils;
|
||||
|
||||
import xiaozhi.common.redis.RedisUtils;
|
||||
import xiaozhi.modules.timbre.dao.TimbreDao;
|
||||
import xiaozhi.modules.timbre.dto.TimbreDataDTO;
|
||||
import xiaozhi.modules.timbre.entity.TimbreEntity;
|
||||
import xiaozhi.modules.timbre.vo.TimbreDetailsVO;
|
||||
import xiaozhi.modules.voiceclone.dao.VoiceCloneDao;
|
||||
import xiaozhi.modules.voiceclone.entity.VoiceCloneEntity;
|
||||
|
||||
@@ -54,4 +58,74 @@ class TimbreServiceImplTest {
|
||||
|
||||
assertNull(service.getDefaultLanguageById("voice-id"));
|
||||
}
|
||||
|
||||
@Test
|
||||
void updateLeavesSortOutOfTheUpdateWhenRequestOmitsIt() {
|
||||
TimbreDao timbreDao = mock(TimbreDao.class);
|
||||
RedisUtils redisUtils = mock(RedisUtils.class);
|
||||
TimbreServiceImpl service = new TimbreServiceImpl(timbreDao, mock(VoiceCloneDao.class), redisUtils);
|
||||
ReflectionTestUtils.setField(service, "baseDao", timbreDao);
|
||||
|
||||
TimbreDataDTO dto = validTimbreData();
|
||||
service.update("voice-id", dto);
|
||||
|
||||
verify(timbreDao, never()).selectById("voice-id");
|
||||
verify(timbreDao).updateById(argThat((TimbreEntity entity) ->
|
||||
"voice-id".equals(entity.getId()) && entity.getSort() == null));
|
||||
verify(redisUtils).delete("timbre:details:voice-id");
|
||||
}
|
||||
|
||||
@Test
|
||||
void updateUsesExplicitSortWithoutLoadingExistingTimbre() {
|
||||
TimbreDao timbreDao = mock(TimbreDao.class);
|
||||
TimbreServiceImpl service = new TimbreServiceImpl(
|
||||
timbreDao, mock(VoiceCloneDao.class), mock(RedisUtils.class));
|
||||
ReflectionTestUtils.setField(service, "baseDao", timbreDao);
|
||||
TimbreDataDTO dto = validTimbreData();
|
||||
dto.setSort(0L);
|
||||
|
||||
service.update("voice-id", dto);
|
||||
|
||||
verify(timbreDao, never()).selectById("voice-id");
|
||||
verify(timbreDao).updateById(argThat((TimbreEntity entity) -> entity.getSort() == 0L));
|
||||
}
|
||||
|
||||
@Test
|
||||
void saveDefaultsOmittedSortToZero() {
|
||||
TimbreDao timbreDao = mock(TimbreDao.class);
|
||||
TimbreServiceImpl service = new TimbreServiceImpl(
|
||||
timbreDao, mock(VoiceCloneDao.class), mock(RedisUtils.class));
|
||||
ReflectionTestUtils.setField(service, "baseDao", timbreDao);
|
||||
|
||||
service.save(validTimbreData());
|
||||
|
||||
verify(timbreDao).insert(argThat((TimbreEntity entity) ->
|
||||
"测试音色".equals(entity.getName()) && entity.getSort() == 0L));
|
||||
}
|
||||
|
||||
@Test
|
||||
void getSupportsLegacyRowsWithNullSort() {
|
||||
TimbreDao timbreDao = mock(TimbreDao.class);
|
||||
RedisUtils redisUtils = mock(RedisUtils.class);
|
||||
TimbreServiceImpl service = new TimbreServiceImpl(
|
||||
timbreDao, mock(VoiceCloneDao.class), redisUtils);
|
||||
ReflectionTestUtils.setField(service, "baseDao", timbreDao);
|
||||
TimbreEntity entity = new TimbreEntity();
|
||||
entity.setId("voice-id");
|
||||
entity.setSort(null);
|
||||
when(timbreDao.selectById("voice-id")).thenReturn(entity);
|
||||
|
||||
TimbreDetailsVO details = service.get("voice-id");
|
||||
|
||||
assertNull(details.getSort());
|
||||
}
|
||||
|
||||
private TimbreDataDTO validTimbreData() {
|
||||
TimbreDataDTO dto = new TimbreDataDTO();
|
||||
dto.setLanguages("中文");
|
||||
dto.setName("测试音色");
|
||||
dto.setTtsModelId("TTS_Test");
|
||||
dto.setTtsVoice("test-voice");
|
||||
return dto;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -69,7 +69,7 @@
|
||||
"build:quickapp-webview-huawei": "uni build -p quickapp-webview-huawei",
|
||||
"build:quickapp-webview-union": "uni build -p quickapp-webview-union",
|
||||
"type-check": "vue-tsc --noEmit",
|
||||
"test:snapshot": "node --test src/pages/agent/components/agentSnapshotUtils.test.mjs src/pages/agent/components/agentSnapshotContracts.test.mjs",
|
||||
"test:snapshot": "node --test src/pages/agent/components/agentSnapshotUtils.test.mjs src/pages/agent/components/agentSnapshotContracts.test.mjs src/pages/agent/components/voicePreviewUtils.test.mjs src/pages/device/deviceTimeUtils.test.mjs",
|
||||
"openapi-ts-request": "openapi-ts",
|
||||
"prepare": "git init && husky",
|
||||
"lint": "eslint",
|
||||
|
||||
@@ -8,6 +8,7 @@ import type {
|
||||
ModelOption,
|
||||
PageData,
|
||||
RoleTemplate,
|
||||
TtsVoice,
|
||||
} from './types'
|
||||
import { http } from '@/http/request/alova'
|
||||
|
||||
@@ -89,7 +90,7 @@ export function deleteAgent(id: string) {
|
||||
|
||||
// 获取TTS音色列表
|
||||
export function getTTSVoices(ttsModelId: string, voiceName: string = '') {
|
||||
return http.Get<{ id: string, name: string }[]>(`/models/${ttsModelId}/voices`, {
|
||||
return http.Get<TtsVoice[]>(`/models/${ttsModelId}/voices`, {
|
||||
params: {
|
||||
voiceName,
|
||||
},
|
||||
@@ -214,7 +215,7 @@ export function updateAgentTags(agentId: string, data) {
|
||||
|
||||
// 获取所有语言
|
||||
export function getAllLanguage(modelId: string) {
|
||||
return http.Get<{ id: string, name: string, languages: string }[]>(`/models/${modelId}/voices`, {
|
||||
return http.Get<TtsVoice[]>(`/models/${modelId}/voices`, {
|
||||
meta: {
|
||||
ignoreAuth: false,
|
||||
toast: false,
|
||||
@@ -225,6 +226,19 @@ export function getAllLanguage(modelId: string) {
|
||||
})
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取克隆音色的临时播放ID
|
||||
* @param cloneId 克隆音色记录ID
|
||||
*/
|
||||
export function getVoiceCloneAudioId(cloneId: string) {
|
||||
return http.Post<string>(`/voiceClone/audio/${cloneId}`, {}, {
|
||||
meta: {
|
||||
ignoreAuth: false,
|
||||
toast: false,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// 获取智能体历史版本列表
|
||||
export function getAgentSnapshots(agentId: string, params: AgentSnapshotPageParams) {
|
||||
return http.Get<PageData<AgentSnapshot>>(`/agent/${agentId}/snapshots`, {
|
||||
|
||||
@@ -114,6 +114,14 @@ export interface CorrectWordFile {
|
||||
wordCount?: number
|
||||
}
|
||||
|
||||
export interface TtsVoice {
|
||||
id: string
|
||||
name: string
|
||||
voiceDemo?: string | null
|
||||
languages?: string | null
|
||||
isClone?: boolean | null
|
||||
}
|
||||
|
||||
// 角色模板数据类型
|
||||
export interface RoleTemplate {
|
||||
id: string
|
||||
|
||||
@@ -7,7 +7,7 @@ export interface Device {
|
||||
id: string
|
||||
userId: string
|
||||
macAddress: string
|
||||
lastConnectedAt: string
|
||||
lastConnectedAtTimestamp: string | null
|
||||
autoUpdate: number
|
||||
board: string
|
||||
alias?: string
|
||||
|
||||
@@ -0,0 +1,42 @@
|
||||
/** @param {Record<string, any>} voice */
|
||||
export function hasVoicePreview(voice) {
|
||||
return Boolean(voice?.isClone || voice?.voiceDemo || voice?.voice_demo)
|
||||
}
|
||||
|
||||
export function createVoicePreviewRequestGate() {
|
||||
let sequence = 0
|
||||
|
||||
return {
|
||||
begin() {
|
||||
sequence += 1
|
||||
return sequence
|
||||
},
|
||||
invalidate() {
|
||||
sequence += 1
|
||||
},
|
||||
isCurrent(requestId) {
|
||||
return requestId === sequence
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* @param {{ id: string, isClone?: boolean, voiceDemo?: string | null }} voice
|
||||
* @param {(cloneId: string) => Promise<string>} getCloneAudioId
|
||||
* @param {string} baseUrl
|
||||
*/
|
||||
export async function resolveVoicePreviewUrl(voice, getCloneAudioId, baseUrl) {
|
||||
if (!voice?.isClone) {
|
||||
return typeof voice?.voiceDemo === 'string' ? voice.voiceDemo : ''
|
||||
}
|
||||
if (!voice.id) {
|
||||
return ''
|
||||
}
|
||||
|
||||
const uuid = await getCloneAudioId(voice.id)
|
||||
if (!uuid) {
|
||||
return ''
|
||||
}
|
||||
|
||||
return `${baseUrl.replace(/\/+$/, '')}/voiceClone/play/${encodeURIComponent(uuid)}`
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
/* eslint-disable test/no-import-node-test -- this zero-dependency gate intentionally uses Node's built-in runner */
|
||||
import assert from 'node:assert/strict'
|
||||
import test from 'node:test'
|
||||
import { createVoicePreviewRequestGate, hasVoicePreview, resolveVoicePreviewUrl } from './voicePreviewUtils.mjs'
|
||||
|
||||
test('keeps normal voice previews on their direct URL', async () => {
|
||||
let cloneRequestCount = 0
|
||||
const url = await resolveVoicePreviewUrl({
|
||||
id: 'normal-voice',
|
||||
isClone: false,
|
||||
voiceDemo: 'https://cdn.example.test/normal.wav',
|
||||
}, async () => {
|
||||
cloneRequestCount += 1
|
||||
return 'unused'
|
||||
}, 'https://api.example.test')
|
||||
|
||||
assert.equal(url, 'https://cdn.example.test/normal.wav')
|
||||
assert.equal(cloneRequestCount, 0)
|
||||
})
|
||||
|
||||
test('uses the clone record id to obtain and construct a temporary play URL', async () => {
|
||||
let requestedCloneId = ''
|
||||
const url = await resolveVoicePreviewUrl({
|
||||
id: 'clone-record-id',
|
||||
isClone: true,
|
||||
voiceDemo: 'provider-speaker-id-must-not-be-played',
|
||||
}, async (cloneId) => {
|
||||
requestedCloneId = cloneId
|
||||
return 'temporary uuid'
|
||||
}, 'https://api.example.test/')
|
||||
|
||||
assert.equal(requestedCloneId, 'clone-record-id')
|
||||
assert.equal(url, 'https://api.example.test/voiceClone/play/temporary%20uuid')
|
||||
})
|
||||
|
||||
test('shows a preview control for cloned voices even without voiceDemo', () => {
|
||||
assert.equal(hasVoicePreview({ id: 'clone-record-id', isClone: true }), true)
|
||||
assert.equal(hasVoicePreview({ id: 'normal-voice', isClone: false, voiceDemo: '' }), false)
|
||||
})
|
||||
|
||||
test('invalidates an older request when the same voice is cancelled and retried', () => {
|
||||
const gate = createVoicePreviewRequestGate()
|
||||
const firstRequest = gate.begin()
|
||||
|
||||
gate.invalidate()
|
||||
const retryRequest = gate.begin()
|
||||
|
||||
assert.equal(gate.isCurrent(firstRequest), false)
|
||||
assert.equal(gate.isCurrent(retryRequest), true)
|
||||
})
|
||||
@@ -1,12 +1,14 @@
|
||||
<script lang="ts" setup>
|
||||
import type { AgentDetail, ModelOption, PluginDefinition, RoleTemplate } from '@/api/agent/types'
|
||||
import type { AgentDetail, ModelOption, PluginDefinition, RoleTemplate, TtsVoice } from '@/api/agent/types'
|
||||
import { computed, nextTick, onMounted, ref, watch } from 'vue'
|
||||
import { getAgentDetail, getAgentTags, getAllLanguage, getModelOptions, getPluginFunctions, getRoleTemplates, updateAgent } from '@/api/agent/agent'
|
||||
import { getAgentDetail, getAgentTags, getAllLanguage, getModelOptions, getPluginFunctions, getRoleTemplates, getVoiceCloneAudioId, updateAgent } from '@/api/agent/agent'
|
||||
import { t } from '@/i18n'
|
||||
import { usePluginStore, useProvider, useSpeedPitch } from '@/store'
|
||||
import { getEnvBaseUrl } from '@/utils'
|
||||
import { toast } from '@/utils/toast'
|
||||
import AgentSnapshotPanel from './components/AgentSnapshotPanel.vue'
|
||||
import { filterTtsVoicesByLanguage, hasUsableTtsVoiceMetadata } from './components/agentSnapshotUtils.mjs'
|
||||
import { createVoicePreviewRequestGate, hasVoicePreview, resolveVoicePreviewUrl } from './components/voicePreviewUtils.mjs'
|
||||
|
||||
defineOptions({
|
||||
name: 'AgentEdit',
|
||||
@@ -84,10 +86,20 @@ const modelOptions = ref<{
|
||||
TTS: [],
|
||||
})
|
||||
|
||||
interface VoiceOption {
|
||||
id?: string
|
||||
value: string
|
||||
name: string
|
||||
voiceDemo?: string | null
|
||||
voice_demo?: string | null
|
||||
isClone: boolean
|
||||
train_status?: number
|
||||
}
|
||||
|
||||
// 音色选项数据
|
||||
const voiceOptions = ref<any[]>([])
|
||||
const voiceOptions = ref<VoiceOption[]>([])
|
||||
// 保存完整的音色信息
|
||||
const voiceDetails = ref<Record<string, any>>({})
|
||||
const voiceDetails = ref<Record<string, TtsVoice>>({})
|
||||
|
||||
// 上报模式选项数据
|
||||
const reportOptions = [
|
||||
@@ -139,6 +151,7 @@ interface SnapshotRestoreContext {
|
||||
// 音频播放相关
|
||||
const audioRef = ref<UniApp.InnerAudioContext | null>(null)
|
||||
const playingVoiceId = ref<string>('')
|
||||
const voicePreviewRequestGate = createVoicePreviewRequestGate()
|
||||
|
||||
// 使用插件store
|
||||
const pluginStore = usePluginStore()
|
||||
@@ -513,8 +526,8 @@ interface TtsSelectionState {
|
||||
languageTouched: boolean
|
||||
voiceTouched: boolean
|
||||
optionsModelId: string
|
||||
voiceOptions: any[]
|
||||
voiceDetails: Record<string, any>
|
||||
voiceOptions: VoiceOption[]
|
||||
voiceDetails: Record<string, TtsVoice>
|
||||
languageOptions: any[]
|
||||
displayNames: {
|
||||
tts: string
|
||||
@@ -565,7 +578,7 @@ function filterVoicesByLanguage(options: VoiceSelectionOptions = {}) {
|
||||
return
|
||||
}
|
||||
|
||||
const allVoices = Object.values(voiceDetails.value) as any[]
|
||||
const allVoices = Object.values(voiceDetails.value)
|
||||
|
||||
// 根据选中的语言筛选音色
|
||||
const filteredVoices = filterTtsVoicesByLanguage(allVoices, selectedTtsLanguage.value)
|
||||
@@ -624,7 +637,7 @@ async function fetchAllLanguag(ttsModelId: string, options: VoiceSelectionOption
|
||||
throw new Error('No TTS voice metadata is available')
|
||||
}
|
||||
// 保存完整的音色信息
|
||||
voiceDetails.value = res.reduce((acc, voice) => {
|
||||
voiceDetails.value = res.reduce<Record<string, TtsVoice>>((acc, voice) => {
|
||||
acc[voice.id] = voice
|
||||
return acc
|
||||
}, {})
|
||||
@@ -863,44 +876,85 @@ function onPickerCancel(type: string) {
|
||||
}
|
||||
|
||||
// 播放音频
|
||||
function playAudio(voiceDemo: string, voiceId: string, event: Event) {
|
||||
async function playAudio(voice: VoiceOption, event: Event) {
|
||||
event.stopPropagation() // 阻止事件冒泡,防止关闭下拉框
|
||||
|
||||
if (!voiceDemo) {
|
||||
if (!hasVoicePreview(voice)) {
|
||||
return
|
||||
}
|
||||
|
||||
// 如果正在播放同一个音频,则停止
|
||||
if (playingVoiceId.value === voiceId) {
|
||||
if (playingVoiceId.value === voice.value) {
|
||||
stopAudio()
|
||||
return
|
||||
}
|
||||
|
||||
// 停止之前的音频
|
||||
stopAudio()
|
||||
const requestId = voicePreviewRequestGate.begin()
|
||||
playingVoiceId.value = voice.value
|
||||
|
||||
// 创建新的音频实例
|
||||
audioRef.value = uni.createInnerAudioContext()
|
||||
audioRef.value.src = voiceDemo
|
||||
playingVoiceId.value = voiceId
|
||||
try {
|
||||
const audioUrl = await resolveVoicePreviewUrl({
|
||||
id: voice.value,
|
||||
isClone: voice.isClone,
|
||||
voiceDemo: voice.voiceDemo || voice.voice_demo,
|
||||
}, getVoiceCloneAudioId, getEnvBaseUrl())
|
||||
|
||||
// 监听播放结束
|
||||
audioRef.value.onEnded(() => {
|
||||
// 用户可能在等待克隆音色临时地址时取消或切换了音色。
|
||||
if (!voicePreviewRequestGate.isCurrent(requestId) || playingVoiceId.value !== voice.value) {
|
||||
return
|
||||
}
|
||||
if (!audioUrl) {
|
||||
toast.error(t('voiceprint.getAudioFailed'))
|
||||
playingVoiceId.value = ''
|
||||
return
|
||||
}
|
||||
|
||||
// 创建新的音频实例
|
||||
const audio = uni.createInnerAudioContext()
|
||||
audioRef.value = audio
|
||||
audio.src = audioUrl
|
||||
|
||||
// 监听播放结束
|
||||
audio.onEnded(() => {
|
||||
if (
|
||||
voicePreviewRequestGate.isCurrent(requestId)
|
||||
&& audioRef.value === audio
|
||||
&& playingVoiceId.value === voice.value
|
||||
) {
|
||||
playingVoiceId.value = ''
|
||||
}
|
||||
})
|
||||
|
||||
// 监听播放错误
|
||||
audio.onError(() => {
|
||||
if (
|
||||
voicePreviewRequestGate.isCurrent(requestId)
|
||||
&& audioRef.value === audio
|
||||
&& playingVoiceId.value === voice.value
|
||||
) {
|
||||
toast.error(t('voiceprint.audioPlayFailed'))
|
||||
playingVoiceId.value = ''
|
||||
}
|
||||
})
|
||||
|
||||
// 播放音频
|
||||
audio.play()
|
||||
}
|
||||
catch (error) {
|
||||
if (!voicePreviewRequestGate.isCurrent(requestId) || playingVoiceId.value !== voice.value) {
|
||||
return
|
||||
}
|
||||
console.error('获取克隆音色试听地址失败:', error)
|
||||
toast.error(t('voiceprint.getAudioFailed'))
|
||||
playingVoiceId.value = ''
|
||||
})
|
||||
|
||||
// 监听播放错误
|
||||
audioRef.value.onError(() => {
|
||||
toast.error('音频播放失败')
|
||||
playingVoiceId.value = ''
|
||||
})
|
||||
|
||||
// 播放音频
|
||||
audioRef.value.play()
|
||||
}
|
||||
}
|
||||
|
||||
// 停止音频
|
||||
function stopAudio() {
|
||||
voicePreviewRequestGate.invalidate()
|
||||
if (audioRef.value) {
|
||||
audioRef.value.stop()
|
||||
audioRef.value.destroy()
|
||||
@@ -1592,10 +1646,10 @@ onMounted(async () => {
|
||||
class="flex items-center justify-between border-b border-[#f5f5f5] p-[32rpx] transition-all active:bg-[#f5f7fb]"
|
||||
@click="onPickerConfirm('voiceprint', voice.value, voice.name)"
|
||||
>
|
||||
<text :class="`flex-1 text-[28rpx] text-[#232338] ${(voice.voiceDemo || voice.voice_demo) ? '' : 'text-center'}`">
|
||||
<text :class="`flex-1 text-[28rpx] text-[#232338] ${hasVoicePreview(voice) ? '' : 'text-center'}`">
|
||||
{{ voice.name }}
|
||||
</text>
|
||||
<view v-if="voice.voiceDemo || voice.voice_demo" class="ml-[20rpx]" @click.stop="playAudio(voice.voiceDemo || voice.voice_demo, voice.value, $event)">
|
||||
<view v-if="hasVoicePreview(voice)" class="ml-[20rpx]" @click.stop="playAudio(voice, $event)">
|
||||
<wd-icon
|
||||
:name="playingVoiceId === voice.value ? 'pause-circle' : 'play-circle'"
|
||||
size="24px"
|
||||
|
||||
@@ -0,0 +1,14 @@
|
||||
/** @param {unknown} timestamp */
|
||||
export function parseDeviceLastConnectedAtTimestamp(timestamp) {
|
||||
if (typeof timestamp !== 'string' || !timestamp.trim()) {
|
||||
return null
|
||||
}
|
||||
|
||||
const milliseconds = Number(timestamp)
|
||||
if (!Number.isFinite(milliseconds)) {
|
||||
return null
|
||||
}
|
||||
|
||||
const date = new Date(milliseconds)
|
||||
return Number.isNaN(date.getTime()) ? null : date
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
/* eslint-disable test/no-import-node-test -- this zero-dependency gate intentionally uses Node's built-in runner */
|
||||
import assert from 'node:assert/strict'
|
||||
import test from 'node:test'
|
||||
import { parseDeviceLastConnectedAtTimestamp } from './deviceTimeUtils.mjs'
|
||||
|
||||
test('parses the backend Long timestamp serialized as a string', () => {
|
||||
const timestamp = '1783689702000'
|
||||
assert.equal(parseDeviceLastConnectedAtTimestamp(timestamp)?.getTime(), Number(timestamp))
|
||||
})
|
||||
|
||||
test('rejects missing and malformed device timestamps', () => {
|
||||
assert.equal(parseDeviceLastConnectedAtTimestamp(null), null)
|
||||
assert.equal(parseDeviceLastConnectedAtTimestamp(''), null)
|
||||
assert.equal(parseDeviceLastConnectedAtTimestamp('not-a-timestamp'), null)
|
||||
})
|
||||
@@ -5,6 +5,7 @@ import { useMessage } from 'wot-design-uni/components/wd-message-box'
|
||||
import { bindDevice, bindDeviceManual, getBindDevices, getFirmwareTypes, unbindDevice, updateDeviceAutoUpdate } from '@/api/device'
|
||||
import { t } from '@/i18n'
|
||||
import { toast } from '@/utils/toast'
|
||||
import { parseDeviceLastConnectedAtTimestamp } from './deviceTimeUtils.mjs'
|
||||
|
||||
defineOptions({
|
||||
name: 'DeviceManage',
|
||||
@@ -131,10 +132,10 @@ function getDeviceTypeName(boardKey: string): string {
|
||||
}
|
||||
|
||||
// 格式化时间
|
||||
function formatTime(timeStr: string) {
|
||||
if (!timeStr)
|
||||
function formatTime(timestamp: string | null) {
|
||||
const date = parseDeviceLastConnectedAtTimestamp(timestamp)
|
||||
if (!date)
|
||||
return t('device.neverConnected')
|
||||
const date = new Date(timeStr)
|
||||
const now = new Date()
|
||||
const diff = now.getTime() - date.getTime()
|
||||
|
||||
@@ -410,7 +411,7 @@ defineExpose({
|
||||
{{ t('device.firmwareVersion') }}:{{ device.appVersion }}
|
||||
</text>
|
||||
<text class="block text-[28rpx] text-[#65686f] leading-[1.4]">
|
||||
{{ t('device.lastConnection') }}:{{ formatTime(device.lastConnectedAt) }}
|
||||
{{ t('device.lastConnection') }}:{{ formatTime(device.lastConnectedAtTimestamp) }}
|
||||
</text>
|
||||
</view>
|
||||
|
||||
|
||||
@@ -27,7 +27,7 @@ export default {
|
||||
getFileList(params, callback) {
|
||||
const queryParams = new URLSearchParams({
|
||||
page: params.page,
|
||||
pageSize: params.pageSize
|
||||
limit: params.pageSize
|
||||
}).toString();
|
||||
|
||||
RequestService.sendRequest()
|
||||
|
||||
@@ -79,6 +79,7 @@ export default {
|
||||
remark: params.remark,
|
||||
referenceAudio: params.referenceAudio,
|
||||
referenceText: params.referenceText,
|
||||
sort: params.sort,
|
||||
ttsModelId: params.ttsModelId,
|
||||
ttsVoice: params.voiceCode,
|
||||
voiceDemo: params.voiceDemo || ''
|
||||
|
||||
@@ -141,23 +141,24 @@
|
||||
<p class="section-desc">{{ $t('addressBookManagement.setPermissionDesc', { count: selectedPermissions.length }) }}</p>
|
||||
</div>
|
||||
<div class="section-actions">
|
||||
<CustomButton size="small" @click="handleCancel">{{ $t('common.cancel') }}</CustomButton>
|
||||
<CustomButton size="small" @click="handleToggleSelectAll">{{ isAllSelected ? $t('addressBookManagement.deselectAll') : $t('addressBookManagement.selectAll') }}</CustomButton>
|
||||
<CustomButton type="confirm" size="small" @click="handleSavePermissions">{{ $t('addressBookManagement.save') }}</CustomButton>
|
||||
<CustomButton size="small" :disabled="permissionsLoading" @click="handleCancel">{{ $t('common.cancel') }}</CustomButton>
|
||||
<CustomButton size="small" :disabled="permissionsLoading" @click="handleToggleSelectAll">{{ isAllSelected ? $t('addressBookManagement.deselectAll') : $t('addressBookManagement.selectAll') }}</CustomButton>
|
||||
<CustomButton type="confirm" size="small" :disabled="permissionsLoading" @click="handleSavePermissions">{{ $t('addressBookManagement.save') }}</CustomButton>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div class="permission-grid">
|
||||
<div v-loading="permissionsLoading" class="permission-grid">
|
||||
<div
|
||||
v-for="device in allDevices"
|
||||
:key="device.id"
|
||||
class="permission-item"
|
||||
:class="{ active: selectedPermissions.includes(device.id) }"
|
||||
:class="{ active: selectedPermissions.includes(device.deviceId) }"
|
||||
>
|
||||
<el-checkbox
|
||||
class="permission-checkbox"
|
||||
:value="selectedPermissions.includes(device.id)"
|
||||
@change="(checked) => handlePermissionToggle(device.id, checked)"
|
||||
:disabled="permissionsLoading"
|
||||
:value="selectedPermissions.includes(device.deviceId)"
|
||||
@change="(checked) => handlePermissionToggle(device.deviceId, checked)"
|
||||
></el-checkbox>
|
||||
<div class="permission-avatar">
|
||||
<img :src="getDeviceAvatar(device.id)" alt="avatar" />
|
||||
@@ -229,7 +230,9 @@ export default {
|
||||
editAgentNameValue: '',
|
||||
editingDeviceId: null,
|
||||
editingDeviceName: '',
|
||||
mqttServiceAvailable: false
|
||||
mqttServiceAvailable: false,
|
||||
permissionRequestSequence: 0,
|
||||
permissionsLoading: false
|
||||
};
|
||||
},
|
||||
created() {
|
||||
@@ -363,18 +366,45 @@ export default {
|
||||
this.loadAddressBookPermissions(device.deviceId);
|
||||
},
|
||||
loadAddressBookPermissions(macAddress) {
|
||||
const requestId = ++this.permissionRequestSequence;
|
||||
this.permissionsLoading = true;
|
||||
this.selectedPermissions = [];
|
||||
this.originalPermissions = [];
|
||||
this.editingDeviceId = null;
|
||||
this.editingDeviceName = '';
|
||||
this.allDevices.forEach(device => {
|
||||
device.addressBookAlias = '';
|
||||
});
|
||||
AddressBookApi.getAddressBookList(macAddress, (res) => {
|
||||
if (
|
||||
requestId !== this.permissionRequestSequence ||
|
||||
this.selectedDevice?.deviceId !== macAddress
|
||||
) {
|
||||
return;
|
||||
}
|
||||
this.permissionsLoading = false;
|
||||
if (res.data?.code === 0) {
|
||||
const permissions = res.data.data || [];
|
||||
const permissionsByTargetMac = new Map(
|
||||
permissions.map(permission => [
|
||||
(permission.targetMac || '').toLowerCase(),
|
||||
permission
|
||||
])
|
||||
);
|
||||
// 设置已选择的权限
|
||||
this.selectedPermissions = permissions
|
||||
.filter(p => p.hasPermission)
|
||||
.map(p => p.targetMac);
|
||||
const permittedTargetMacs = new Set(
|
||||
permissions
|
||||
.filter(p => p.hasPermission)
|
||||
.map(p => (p.targetMac || '').toLowerCase())
|
||||
);
|
||||
this.selectedPermissions = this.allDevices
|
||||
.filter(device => permittedTargetMacs.has((device.deviceId || '').toLowerCase()))
|
||||
.map(device => device.deviceId);
|
||||
// 保存初始权限状态(用于对比变更)
|
||||
this.originalPermissions = [...this.selectedPermissions];
|
||||
// 更新设备的通讯录别名
|
||||
this.allDevices.forEach(device => {
|
||||
const addrBook = permissions.find(p => p.targetMac === device.deviceId);
|
||||
const addrBook = permissionsByTargetMac.get((device.deviceId || '').toLowerCase());
|
||||
if (addrBook) {
|
||||
device.addressBookAlias = addrBook.alias || '';
|
||||
} else {
|
||||
@@ -385,6 +415,7 @@ export default {
|
||||
});
|
||||
},
|
||||
handleStartEditPermission(device) {
|
||||
if (this.permissionsLoading) return;
|
||||
this.editingDeviceId = device.id;
|
||||
this.editingDeviceName = device.addressBookAlias || device.name;
|
||||
this.$nextTick(() => {
|
||||
@@ -412,13 +443,14 @@ export default {
|
||||
this.editingDeviceId = null;
|
||||
this.editingDeviceName = '';
|
||||
},
|
||||
handlePermissionToggle(deviceId, checked) {
|
||||
handlePermissionToggle(targetMac, checked) {
|
||||
if (this.permissionsLoading) return;
|
||||
if (checked) {
|
||||
if (!this.selectedPermissions.includes(deviceId)) {
|
||||
this.selectedPermissions.push(deviceId);
|
||||
if (!this.selectedPermissions.includes(targetMac)) {
|
||||
this.selectedPermissions.push(targetMac);
|
||||
}
|
||||
} else {
|
||||
const index = this.selectedPermissions.indexOf(deviceId);
|
||||
const index = this.selectedPermissions.indexOf(targetMac);
|
||||
if (index > -1) {
|
||||
this.selectedPermissions.splice(index, 1);
|
||||
}
|
||||
@@ -428,21 +460,22 @@ export default {
|
||||
if (this.isAllSelected) {
|
||||
this.selectedPermissions = [];
|
||||
} else {
|
||||
this.selectedPermissions = this.allDevices.map(d => d.id);
|
||||
this.selectedPermissions = this.allDevices.map(d => d.deviceId);
|
||||
}
|
||||
},
|
||||
handleCancel() {
|
||||
this.selectedPermissions = [];
|
||||
},
|
||||
handleSavePermissions() {
|
||||
if (this.permissionsLoading) return;
|
||||
const promises = this.allDevices
|
||||
.filter(device => {
|
||||
const isNowSelected = this.selectedPermissions.includes(device.id);
|
||||
const wasOriginallySelected = this.originalPermissions.includes(device.id);
|
||||
const isNowSelected = this.selectedPermissions.includes(device.deviceId);
|
||||
const wasOriginallySelected = this.originalPermissions.includes(device.deviceId);
|
||||
return isNowSelected !== wasOriginallySelected;
|
||||
})
|
||||
.map(device => {
|
||||
const hasPermission = this.selectedPermissions.includes(device.id);
|
||||
const hasPermission = this.selectedPermissions.includes(device.deviceId);
|
||||
return new Promise((resolve) => {
|
||||
AddressBookApi.updatePermission({
|
||||
macAddress: this.selectedDevice.deviceId,
|
||||
|
||||
@@ -0,0 +1,63 @@
|
||||
import assert from 'node:assert/strict';
|
||||
import { readFile } from 'node:fs/promises';
|
||||
import test from 'node:test';
|
||||
|
||||
const addressBookSource = await readFile(
|
||||
new URL('../src/views/AddressBookManagement.vue', import.meta.url),
|
||||
'utf8',
|
||||
);
|
||||
const correctWordApiSource = await readFile(
|
||||
new URL('../src/apis/module/correctWord.js', import.meta.url),
|
||||
'utf8',
|
||||
);
|
||||
|
||||
test('address-book permission state consistently uses the target device MAC', () => {
|
||||
assert.match(
|
||||
addressBookSource,
|
||||
/:value="selectedPermissions\.includes\(device\.deviceId\)"/,
|
||||
);
|
||||
assert.match(
|
||||
addressBookSource,
|
||||
/@change="\(checked\) => handlePermissionToggle\(device\.deviceId, checked\)"/,
|
||||
);
|
||||
assert.match(
|
||||
addressBookSource,
|
||||
/this\.selectedPermissions = this\.allDevices\.map\(d => d\.deviceId\)/,
|
||||
);
|
||||
assert.match(
|
||||
addressBookSource,
|
||||
/this\.originalPermissions\.includes\(device\.deviceId\)/,
|
||||
);
|
||||
assert.doesNotMatch(
|
||||
addressBookSource,
|
||||
/selectedPermissions\.includes\(device\.id\)/,
|
||||
);
|
||||
assert.doesNotMatch(
|
||||
addressBookSource,
|
||||
/originalPermissions\.includes\(device\.id\)/,
|
||||
);
|
||||
assert.match(
|
||||
addressBookSource,
|
||||
/requestId !== this\.permissionRequestSequence/,
|
||||
);
|
||||
assert.match(
|
||||
addressBookSource,
|
||||
/this\.selectedDevice\?\.deviceId !== macAddress/,
|
||||
);
|
||||
assert.match(
|
||||
addressBookSource,
|
||||
/this\.permissionsLoading = true;\s*this\.selectedPermissions = \[\];\s*this\.originalPermissions = \[\];/,
|
||||
);
|
||||
assert.match(
|
||||
addressBookSource,
|
||||
/handleSavePermissions\(\) \{\s*if \(this\.permissionsLoading\) return;/,
|
||||
);
|
||||
});
|
||||
|
||||
test('correct-word pagination maps the UI page size to the backend limit query', () => {
|
||||
assert.match(
|
||||
correctWordApiSource,
|
||||
/new URLSearchParams\(\{\s*page: params\.page,\s*limit: params\.pageSize\s*\}\)/,
|
||||
);
|
||||
assert.doesNotMatch(correctWordApiSource, /pageSize: params\.pageSize/);
|
||||
});
|
||||
@@ -0,0 +1,17 @@
|
||||
import assert from 'node:assert/strict';
|
||||
import { readFile } from 'node:fs/promises';
|
||||
import test from 'node:test';
|
||||
|
||||
const timbreApiSource = await readFile(
|
||||
new URL('../src/apis/module/timbre.js', import.meta.url),
|
||||
'utf8',
|
||||
);
|
||||
|
||||
test('timbre update sends the current sort value to the backend', () => {
|
||||
const updateStart = timbreApiSource.indexOf('updateVoice(params, callback)');
|
||||
|
||||
assert.notEqual(updateStart, -1);
|
||||
const updateSource = timbreApiSource.slice(updateStart);
|
||||
assert.match(updateSource, /\.method\('PUT'\)/);
|
||||
assert.match(updateSource, /sort:\s*params\.sort/);
|
||||
});
|
||||
+34
-143
@@ -3,14 +3,13 @@ import uuid
|
||||
import signal
|
||||
import asyncio
|
||||
from aioconsole import ainput
|
||||
from config.config_loader import load_config
|
||||
from config.settings 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.xiaozhi_server_facade import XiaozhiServerFacade
|
||||
from core.websocket_server import WebSocketServer
|
||||
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()
|
||||
@@ -69,49 +68,12 @@ async def main():
|
||||
gc_manager = get_gc_manager(interval_seconds=300)
|
||||
await gc_manager.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
|
||||
# 启动 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())
|
||||
|
||||
read_config_from_api = config.get("read_config_from_api", False)
|
||||
port = int(config["server"].get("http_port", 8003))
|
||||
@@ -138,51 +100,24 @@ async def main():
|
||||
logger.bind(tag=TAG).error("mcp接入点不符合规范")
|
||||
config["mcp_endpoint"] = "你的接入点 websocket地址"
|
||||
|
||||
# 显示协议连接信息
|
||||
connection_info = xiaozhi_server.get_connection_info()
|
||||
# 获取WebSocket配置,使用安全的默认值
|
||||
websocket_port = 8000
|
||||
server_config = config.get("server", {})
|
||||
if isinstance(server_config, dict):
|
||||
websocket_port = int(server_config.get("port", 8000))
|
||||
|
||||
# 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地址是\tws://{}:{}/xiaozhi/v1/",
|
||||
get_local_ip(),
|
||||
websocket_port,
|
||||
)
|
||||
|
||||
# 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(
|
||||
"=======上面的地址是websocket协议地址,请勿用浏览器访问======="
|
||||
)
|
||||
logger.bind(tag=TAG).info(
|
||||
"如想测试websocket请启动digital-human模块,打开浏览器交互测试"
|
||||
)
|
||||
logger.bind(tag=TAG).info(
|
||||
"=============================================================\n"
|
||||
)
|
||||
@@ -192,66 +127,22 @@ 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管理器
|
||||
try:
|
||||
await gc_manager.stop()
|
||||
except Exception as e:
|
||||
shutdown_errors.append(("GC管理器", e))
|
||||
logger.bind(tag=TAG).error(f"停止GC管理器失败: {e}")
|
||||
await gc_manager.stop()
|
||||
|
||||
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()
|
||||
await asyncio.gather(stdin_task, return_exceptions=True)
|
||||
|
||||
ws_task.cancel()
|
||||
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}"
|
||||
)
|
||||
ota_task.cancel()
|
||||
|
||||
# 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}")
|
||||
# 等待任务终止(必须加超时)
|
||||
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,
|
||||
)
|
||||
print("服务器已关闭,程序退出。")
|
||||
if shutdown_errors:
|
||||
details = ", ".join(
|
||||
f"{owner}: {error}" for owner, error in shutdown_errors
|
||||
)
|
||||
raise RuntimeError(f"服务器清理失败: {details}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -35,75 +35,12 @@ 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,28 +80,6 @@ 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请求并处理响应"""
|
||||
@@ -288,7 +266,3 @@ def init_service(config):
|
||||
|
||||
def manage_api_http_safe_close():
|
||||
ManageApiClient.safe_close()
|
||||
|
||||
|
||||
async def manage_api_http_close():
|
||||
await ManageApiClient.close_all_clients()
|
||||
|
||||
@@ -1,323 +0,0 @@
|
||||
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,6 +1,8 @@
|
||||
import json
|
||||
import time
|
||||
import base64
|
||||
import hashlib
|
||||
import hmac
|
||||
import os
|
||||
import re
|
||||
import glob
|
||||
@@ -9,11 +11,6 @@ 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__
|
||||
@@ -105,6 +102,26 @@ 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地址
|
||||
|
||||
@@ -214,32 +231,11 @@ class OTAHandler(BaseHandler):
|
||||
},
|
||||
}
|
||||
|
||||
# ========== 协议下发逻辑 ==========
|
||||
# 按照原版逻辑:总是下发 WebSocket,如果启用了 MQTT 则额外下发 MQTT 和 UDP
|
||||
# 这样设备有回退能力:如果 MQTT 连接失败,还可以使用 WebSocket
|
||||
|
||||
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
|
||||
)
|
||||
# existing mqtt/websocket logic (unchanged)
|
||||
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():
|
||||
if mqtt_gateway_endpoint: # 如果配置了非空字符串
|
||||
# 尝试从请求数据中获取设备型号(已解析 above)
|
||||
try:
|
||||
group_id = f"GID_{device_model}".replace(":", "_").replace(" ", "_")
|
||||
except Exception as e:
|
||||
@@ -250,116 +246,56 @@ class OTAHandler(BaseHandler):
|
||||
mqtt_client_id = f"{group_id}@@@{mac_address_safe}@@@{mac_address_safe}"
|
||||
|
||||
# 构建用户数据
|
||||
user_data = {"ip": local_ip}
|
||||
user_data = {"ip": "unknown"}
|
||||
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()
|
||||
|
||||
# 生成密码
|
||||
mqtt_password = ""
|
||||
password = ""
|
||||
signature_key = server_config.get("mqtt_signature_key", "")
|
||||
if signature_key:
|
||||
mqtt_password = generate_password_signature(
|
||||
password = self.generate_password_signature(
|
||||
mqtt_client_id + "|" + username, signature_key
|
||||
)
|
||||
if not mqtt_password:
|
||||
mqtt_password = ""
|
||||
if not password:
|
||||
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": mqtt_password,
|
||||
"password": password,
|
||||
"publish_topic": "device-server",
|
||||
"subscribe_topic": f"devices/p2p/{mac_address_safe}",
|
||||
}
|
||||
self.logger.bind(tag=TAG).info(f"为设备 {device_id} 下发MQTT网关配置")
|
||||
|
||||
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)}")
|
||||
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配置"
|
||||
)
|
||||
|
||||
# Now check firmware files for updates
|
||||
try:
|
||||
|
||||
@@ -10,164 +10,6 @@ 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:
|
||||
"""
|
||||
统一授权认证管理器
|
||||
|
||||
@@ -1,116 +0,0 @@
|
||||
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
|
||||
@@ -1,92 +0,0 @@
|
||||
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
|
||||
@@ -1,64 +0,0 @@
|
||||
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
|
||||
@@ -1,103 +0,0 @@
|
||||
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
|
||||
@@ -1,76 +0,0 @@
|
||||
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
|
||||
@@ -1,64 +0,0 @@
|
||||
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
|
||||
@@ -1,268 +0,0 @@
|
||||
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()
|
||||
}
|
||||
@@ -1,73 +0,0 @@
|
||||
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
|
||||
@@ -1,505 +0,0 @@
|
||||
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__()
|
||||
@@ -350,16 +350,4 @@ async def send_display_message(conn: "ConnectionHandler", text):
|
||||
"text": text,
|
||||
"session_id": conn.session_id
|
||||
}
|
||||
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))
|
||||
await conn.websocket.send(json.dumps(message))
|
||||
|
||||
@@ -1,9 +1,6 @@
|
||||
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
|
||||
|
||||
@@ -11,53 +8,11 @@ TAG = __name__
|
||||
|
||||
|
||||
class SimpleHttpServer:
|
||||
def __init__(self, config: dict, management_owner=None):
|
||||
def __init__(self, config: dict):
|
||||
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地址
|
||||
@@ -78,28 +33,14 @@ 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:
|
||||
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)
|
||||
app = web.Application()
|
||||
|
||||
if not read_config_from_api:
|
||||
# 如果没有开启智控台,只是单模块运行,就需要再添加简单OTA接口,用于下发websocket接口
|
||||
@@ -133,61 +74,19 @@ 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()
|
||||
self._started_event.set()
|
||||
await self._stop_event.wait()
|
||||
|
||||
# 保持服务运行
|
||||
while True:
|
||||
await asyncio.sleep(3600) # 每隔 1 小时检查一次
|
||||
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()
|
||||
|
||||
@@ -1,26 +0,0 @@
|
||||
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
|
||||
@@ -1,132 +0,0 @@
|
||||
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)
|
||||
@@ -1,439 +0,0 @@
|
||||
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
|
||||
)
|
||||
@@ -1,578 +0,0 @@
|
||||
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, "再见,期待下次相遇")
|
||||
@@ -1,55 +0,0 @@
|
||||
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
@@ -1,46 +0,0 @@
|
||||
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
|
||||
@@ -1,351 +0,0 @@
|
||||
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)
|
||||
@@ -1,124 +0,0 @@
|
||||
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}")
|
||||
@@ -1,433 +0,0 @@
|
||||
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)
|
||||
@@ -1,96 +0,0 @@
|
||||
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}")
|
||||
@@ -1,127 +0,0 @@
|
||||
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]
|
||||
@@ -1,40 +0,0 @@
|
||||
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
|
||||
@@ -1,301 +0,0 @@
|
||||
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}")
|
||||
@@ -1,177 +0,0 @@
|
||||
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
|
||||
@@ -1,54 +0,0 @@
|
||||
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}")
|
||||
@@ -1,48 +0,0 @@
|
||||
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
|
||||
@@ -1,52 +0,0 @@
|
||||
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
|
||||
@@ -1,649 +0,0 @@
|
||||
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}")
|
||||
@@ -1,651 +0,0 @@
|
||||
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()
|
||||
@@ -169,10 +169,7 @@ class ASRProvider(ASRProviderBase):
|
||||
|
||||
self.is_processing = True
|
||||
self.server_ready = False # 重置服务器准备状态
|
||||
session_id = getattr(conn, "session_id", None)
|
||||
self.forward_task = self._create_session_task(
|
||||
conn, self._forward_results(conn, session_id)
|
||||
)
|
||||
self.forward_task = asyncio.create_task(self._forward_results(conn))
|
||||
|
||||
# 发送开始请求
|
||||
start_request = {
|
||||
@@ -196,13 +193,10 @@ 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", session_id=None):
|
||||
async def _forward_results(self, conn: "ConnectionHandler"):
|
||||
"""转发识别结果"""
|
||||
try:
|
||||
while (
|
||||
not conn.stop_event.is_set()
|
||||
and self._session_is_current(conn, session_id)
|
||||
):
|
||||
while not conn.stop_event.is_set():
|
||||
# 获取当前连接的音频数据
|
||||
audio_data = conn.asr_audio
|
||||
try:
|
||||
@@ -282,7 +276,7 @@ class ASRProvider(ASRProviderBase):
|
||||
finally:
|
||||
# 清理连接的音频缓存
|
||||
await self._cleanup()
|
||||
self._reset_audio_if_current(conn, session_id)
|
||||
conn.reset_audio_states()
|
||||
|
||||
async def _send_stop_request(self):
|
||||
"""发送停止识别请求(不关闭连接)"""
|
||||
@@ -314,21 +308,6 @@ 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:
|
||||
@@ -340,10 +319,8 @@ class ASRProvider(ASRProviderBase):
|
||||
finally:
|
||||
self.asr_ws = 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
|
||||
# 清理任务引用
|
||||
self.forward_task = None
|
||||
|
||||
logger.bind(tag=TAG).debug("ASR会话清理完成")
|
||||
|
||||
|
||||
@@ -103,10 +103,7 @@ class ASRProvider(ASRProviderBase):
|
||||
logger.bind(tag=TAG).debug("WebSocket连接建立成功")
|
||||
|
||||
self.server_ready = False
|
||||
session_id = getattr(conn, "session_id", None)
|
||||
self.forward_task = self._create_session_task(
|
||||
conn, self._forward_results(conn, session_id)
|
||||
)
|
||||
self.forward_task = asyncio.create_task(self._forward_results(conn))
|
||||
|
||||
# 发送run-task指令
|
||||
run_task_msg = self._build_run_task_message()
|
||||
@@ -157,13 +154,10 @@ class ASRProvider(ASRProviderBase):
|
||||
|
||||
return message
|
||||
|
||||
async def _forward_results(self, conn: "ConnectionHandler", session_id=None):
|
||||
async def _forward_results(self, conn: "ConnectionHandler"):
|
||||
"""转发识别结果"""
|
||||
try:
|
||||
while (
|
||||
not conn.stop_event.is_set()
|
||||
and self._session_is_current(conn, session_id)
|
||||
):
|
||||
while not conn.stop_event.is_set():
|
||||
# 获取当前连接的音频数据
|
||||
audio_data = conn.asr_audio
|
||||
try:
|
||||
@@ -249,7 +243,7 @@ class ASRProvider(ASRProviderBase):
|
||||
finally:
|
||||
# 清理连接的音频缓存
|
||||
await self._cleanup()
|
||||
self._reset_audio_if_current(conn, session_id)
|
||||
conn.reset_audio_states()
|
||||
|
||||
async def _send_stop_request(self):
|
||||
"""发送停止请求(用于手动模式停止录音)"""
|
||||
@@ -291,21 +285,6 @@ 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:
|
||||
@@ -322,8 +301,8 @@ class ASRProvider(ASRProviderBase):
|
||||
finally:
|
||||
self.asr_ws = None
|
||||
|
||||
if self.forward_task is forward_task:
|
||||
self.forward_task = None
|
||||
# 清理任务引用
|
||||
self.forward_task = None
|
||||
self.task_id = None
|
||||
|
||||
logger.bind(tag=TAG).debug("ASR会话清理完成")
|
||||
|
||||
@@ -32,37 +32,8 @@ 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
|
||||
)
|
||||
@@ -101,26 +72,18 @@ class ASRProviderBase(ABC):
|
||||
return
|
||||
|
||||
# 自动模式下通过VAD检测到语音停止时触发识别
|
||||
interface_type = getattr(
|
||||
self,
|
||||
"interface_type",
|
||||
getattr(getattr(conn, "asr", None), "interface_type", None),
|
||||
)
|
||||
if interface_type != InterfaceType.STREAM and conn.client_voice_stop:
|
||||
if conn.asr.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])
|
||||
reset_audio_states = getattr(conn, "reset_audio_states", None)
|
||||
if callable(reset_audio_states):
|
||||
reset_audio_states()
|
||||
conn.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直接使用
|
||||
@@ -133,11 +96,13 @@ class ASRProviderBase(ABC):
|
||||
wav_data = self._pcm_to_wav(combined_pcm_data)
|
||||
|
||||
# 定义ASR任务
|
||||
asr_task = self.speech_to_text_wrapper(asr_audio_task, session_id)
|
||||
asr_task = self.speech_to_text_wrapper(
|
||||
asr_audio_task, conn.session_id
|
||||
)
|
||||
|
||||
if conn.voiceprint_provider and wav_data:
|
||||
voiceprint_task = conn.voiceprint_provider.identify_speaker(
|
||||
wav_data, session_id
|
||||
wav_data, conn.session_id
|
||||
)
|
||||
# 并发等待两个结果
|
||||
asr_result, voiceprint_result = await asyncio.gather(
|
||||
@@ -147,12 +112,6 @@ 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}")
|
||||
@@ -206,13 +165,9 @@ class ASRProviderBase(ABC):
|
||||
|
||||
if text_len > 0:
|
||||
audio_snapshot = asr_audio_task.copy()
|
||||
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)
|
||||
enqueue_asr_report(conn, enhanced_text, audio_snapshot)
|
||||
# 使用自定义模块进行上报
|
||||
await startToChat(conn, enhanced_text)
|
||||
except Exception as e:
|
||||
logger.bind(tag=TAG).error(f"处理语音停止失败: {e}")
|
||||
import traceback
|
||||
|
||||
@@ -117,10 +117,7 @@ class ASRProvider(ASRProviderBase):
|
||||
raise e
|
||||
|
||||
# 启动接收ASR结果的异步任务
|
||||
session_id = getattr(conn, "session_id", None)
|
||||
self.forward_task = self._create_session_task(
|
||||
conn, self._forward_asr_results(conn, session_id)
|
||||
)
|
||||
self.forward_task = asyncio.create_task(self._forward_asr_results(conn))
|
||||
|
||||
# 发送缓存的音频数据
|
||||
if conn.asr_audio and len(conn.asr_audio) > 0:
|
||||
@@ -159,13 +156,9 @@ class ASRProvider(ASRProviderBase):
|
||||
except Exception as e:
|
||||
logger.bind(tag=TAG).info(f"发送音频数据时发生错误: {e}")
|
||||
|
||||
async def _forward_asr_results(self, conn: "ConnectionHandler", session_id=None):
|
||||
async def _forward_asr_results(self, conn: "ConnectionHandler"):
|
||||
try:
|
||||
while (
|
||||
self.asr_ws
|
||||
and not conn.stop_event.is_set()
|
||||
and self._session_is_current(conn, session_id)
|
||||
):
|
||||
while self.asr_ws and not conn.stop_event.is_set():
|
||||
# 获取当前连接的音频数据
|
||||
audio_data = conn.asr_audio
|
||||
try:
|
||||
@@ -256,46 +249,20 @@ class ASRProvider(ASRProviderBase):
|
||||
if hasattr(e, "__cause__") and e.__cause__:
|
||||
logger.bind(tag=TAG).error(f"错误原因: {str(e.__cause__)}")
|
||||
finally:
|
||||
await self._cleanup()
|
||||
if self.asr_ws:
|
||||
await self.asr_ws.close()
|
||||
self.asr_ws = None
|
||||
self.is_processing = False
|
||||
self._is_stopping = False
|
||||
# 重置所有音频相关状态
|
||||
self._reset_audio_if_current(conn, session_id)
|
||||
conn.reset_audio_states()
|
||||
|
||||
def stop_ws_connection(self):
|
||||
# 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
|
||||
asyncio.create_task(self.asr_ws.close())
|
||||
self.asr_ws = None
|
||||
self.is_processing = False
|
||||
self._is_stopping = False
|
||||
|
||||
async def _send_stop_request(self):
|
||||
"""发送最后一个音频帧以通知服务器结束"""
|
||||
@@ -450,4 +417,14 @@ class ASRProvider(ASRProviderBase):
|
||||
|
||||
async def close(self):
|
||||
"""资源清理方法"""
|
||||
await self._cleanup()
|
||||
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
|
||||
|
||||
@@ -90,11 +90,8 @@ 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 | 结果: {recognized_content}"
|
||||
f"语音识别耗时: {time.time() - start_time:.3f}s | 结果: {text['content']}"
|
||||
)
|
||||
|
||||
return text, artifacts.file_path
|
||||
|
||||
@@ -1,522 +0,0 @@
|
||||
"""
|
||||
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
|
||||
@@ -1,119 +0,0 @@
|
||||
"""
|
||||
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,10 +141,7 @@ class ASRProvider(ASRProviderBase):
|
||||
|
||||
logger.bind(tag=TAG).info("ASR WebSocket连接已建立")
|
||||
self.server_ready = False
|
||||
session_id = getattr(conn, "session_id", None)
|
||||
self.forward_task = self._create_session_task(
|
||||
conn, self._forward_results(conn, session_id)
|
||||
)
|
||||
self.forward_task = asyncio.create_task(self._forward_results(conn))
|
||||
|
||||
# 发送首帧音频
|
||||
if conn.asr_audio and len(conn.asr_audio) > 0:
|
||||
@@ -188,13 +185,10 @@ class ASRProvider(ASRProviderBase):
|
||||
|
||||
await self.asr_ws.send(json.dumps(frame_data, ensure_ascii=False))
|
||||
|
||||
async def _forward_results(self, conn: "ConnectionHandler", session_id=None):
|
||||
async def _forward_results(self, conn: "ConnectionHandler"):
|
||||
"""转发识别结果"""
|
||||
try:
|
||||
while (
|
||||
not conn.stop_event.is_set()
|
||||
and self._session_is_current(conn, session_id)
|
||||
):
|
||||
while not conn.stop_event.is_set():
|
||||
try:
|
||||
response = await asyncio.wait_for(self.asr_ws.recv(), timeout=60)
|
||||
result = json.loads(response)
|
||||
@@ -253,7 +247,7 @@ class ASRProvider(ASRProviderBase):
|
||||
finally:
|
||||
# 清理连接资源
|
||||
await self._cleanup()
|
||||
self._reset_audio_if_current(conn, session_id)
|
||||
conn.reset_audio_states()
|
||||
|
||||
async def handle_voice_stop(
|
||||
self, conn: "ConnectionHandler", asr_audio_task: List[bytes]
|
||||
@@ -278,8 +272,9 @@ class ASRProvider(ASRProviderBase):
|
||||
logger.bind(tag=TAG).debug(f"异常详情: {traceback.format_exc()}")
|
||||
|
||||
def stop_ws_connection(self):
|
||||
# The forward task owns the WebSocket and closes it from _cleanup().
|
||||
# Scheduling an untracked close here races with that cleanup path.
|
||||
if self.asr_ws:
|
||||
asyncio.create_task(self.asr_ws.close())
|
||||
self.asr_ws = None
|
||||
self.is_processing = False
|
||||
|
||||
async def _send_stop_request(self):
|
||||
@@ -304,21 +299,6 @@ 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:
|
||||
@@ -330,8 +310,8 @@ class ASRProvider(ASRProviderBase):
|
||||
finally:
|
||||
self.asr_ws = None
|
||||
|
||||
if self.forward_task is forward_task:
|
||||
self.forward_task = None
|
||||
# 清理任务引用
|
||||
self.forward_task = None
|
||||
|
||||
logger.bind(tag=TAG).debug("ASR会话清理完成")
|
||||
|
||||
@@ -343,4 +323,15 @@ class ASRProvider(ASRProviderBase):
|
||||
|
||||
async def close(self):
|
||||
"""资源清理方法"""
|
||||
await self._cleanup()
|
||||
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
|
||||
|
||||
|
||||
@@ -127,16 +127,7 @@ class DeviceIoTExecutor(ToolExecutor):
|
||||
send_message = json.dumps(
|
||||
{"type": "iot", "commands": [command]}
|
||||
)
|
||||
|
||||
# 使用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("无法找到可用的传输层接口")
|
||||
await self.conn.websocket.send(send_message)
|
||||
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 = 10000
|
||||
self.next_id = 1
|
||||
self.lock = asyncio.Lock()
|
||||
self._cached_available_tools = None # Cache for get_available_tools
|
||||
|
||||
@@ -91,12 +91,3 @@ 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 core.utils.util import get_vision_url
|
||||
from concurrent.futures import Future
|
||||
from core.utils.util import get_vision_url, sanitize_tool_name
|
||||
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,45 +15,108 @@ if TYPE_CHECKING:
|
||||
TAG = __name__
|
||||
logger = setup_logging()
|
||||
|
||||
async def send_mcp_message(
|
||||
conn: "ConnectionHandler",
|
||||
payload: dict,
|
||||
transport=None,
|
||||
*,
|
||||
raise_on_error: bool = False,
|
||||
):
|
||||
|
||||
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):
|
||||
"""Helper to send MCP messages, encapsulating common logic."""
|
||||
features = getattr(conn, "features", {}) or {}
|
||||
if not features.get("mcp"):
|
||||
if not conn.features.get("mcp"):
|
||||
logger.bind(tag=TAG).warning("客户端不支持MCP,无法发送MCP消息")
|
||||
return
|
||||
|
||||
message = json.dumps({"type": "mcp", "payload": payload})
|
||||
|
||||
try:
|
||||
# 优先使用传入的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("无法找到可用的传输层接口")
|
||||
await conn.websocket.send(message)
|
||||
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,
|
||||
transport=None,
|
||||
conn: "ConnectionHandler", mcp_client: MCPClient, payload: dict
|
||||
):
|
||||
"""处理MCP消息,包括初始化、工具列表和工具调用响应等"""
|
||||
logger.bind(tag=TAG).debug(f"处理MCP消息: {str(payload)[:100]}")
|
||||
@@ -87,7 +150,7 @@ async def handle_mcp_message(
|
||||
|
||||
await asyncio.sleep(1)
|
||||
logger.bind(tag=TAG).debug("初始化完成,开始请求MCP工具列表")
|
||||
await send_mcp_tools_list_request(conn, transport)
|
||||
await send_mcp_tools_list_request(conn)
|
||||
|
||||
return
|
||||
|
||||
@@ -144,7 +207,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, transport)
|
||||
await send_mcp_tools_list_continue_request(conn, next_cursor)
|
||||
else:
|
||||
await mcp_client.set_ready(True)
|
||||
logger.bind(tag=TAG).debug("所有工具已获取,MCP客户端准备就绪")
|
||||
@@ -172,9 +235,7 @@ async def handle_mcp_message(
|
||||
)
|
||||
|
||||
|
||||
async def send_mcp_initialize_message(
|
||||
conn: "ConnectionHandler", transport=None
|
||||
):
|
||||
async def send_mcp_initialize_message(conn: "ConnectionHandler"):
|
||||
"""发送MCP初始化消息"""
|
||||
|
||||
vision_url = get_vision_url(conn.config)
|
||||
@@ -206,12 +267,10 @@ async def send_mcp_initialize_message(
|
||||
},
|
||||
}
|
||||
logger.bind(tag=TAG).debug("发送MCP初始化消息")
|
||||
await send_mcp_message(conn, payload, transport)
|
||||
await send_mcp_message(conn, payload)
|
||||
|
||||
|
||||
async def send_mcp_tools_list_request(
|
||||
conn: "ConnectionHandler", transport=None
|
||||
):
|
||||
async def send_mcp_tools_list_request(conn: "ConnectionHandler"):
|
||||
"""发送MCP工具列表请求"""
|
||||
payload = {
|
||||
"jsonrpc": "2.0",
|
||||
@@ -219,12 +278,10 @@ async def send_mcp_tools_list_request(
|
||||
"method": "tools/list",
|
||||
}
|
||||
logger.bind(tag=TAG).debug("发送MCP工具列表请求")
|
||||
await send_mcp_message(conn, payload, transport)
|
||||
await send_mcp_message(conn, payload)
|
||||
|
||||
|
||||
async def send_mcp_tools_list_continue_request(
|
||||
conn: "ConnectionHandler", cursor: str, transport=None
|
||||
):
|
||||
async def send_mcp_tools_list_continue_request(conn: "ConnectionHandler", cursor: str):
|
||||
"""发送带有cursor的MCP工具列表请求"""
|
||||
payload = {
|
||||
"jsonrpc": "2.0",
|
||||
@@ -233,7 +290,7 @@ async def send_mcp_tools_list_continue_request(
|
||||
"params": {"cursor": cursor},
|
||||
}
|
||||
logger.bind(tag=TAG).info(f"发送带cursor的MCP工具列表请求: {cursor}")
|
||||
await send_mcp_message(conn, payload, transport)
|
||||
await send_mcp_message(conn, payload)
|
||||
|
||||
|
||||
async def call_mcp_tool(
|
||||
@@ -242,7 +299,6 @@ async def call_mcp_tool(
|
||||
tool_name: str,
|
||||
args: str = "{}",
|
||||
timeout: int = 30,
|
||||
return_raw: bool = False,
|
||||
):
|
||||
"""
|
||||
调用指定的工具,并等待响应
|
||||
@@ -303,7 +359,6 @@ 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
|
||||
@@ -317,11 +372,7 @@ async def call_mcp_tool(
|
||||
}
|
||||
|
||||
logger.bind(tag=TAG).info(f"发送客户端mcp工具调用请求: {actual_name},参数: {args}")
|
||||
try:
|
||||
await send_mcp_message(conn, payload, raise_on_error=True)
|
||||
except Exception:
|
||||
await mcp_client.cleanup_call_result(tool_call_id)
|
||||
raise
|
||||
await send_mcp_message(conn, payload)
|
||||
|
||||
try:
|
||||
# Wait for response or timeout
|
||||
@@ -337,21 +388,13 @@ 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,7 +22,6 @@ 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
|
||||
@@ -108,18 +107,6 @@ 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,8 +16,6 @@ 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)
|
||||
|
||||
@@ -25,9 +23,7 @@ async def connect_mcp_endpoint(mcp_endpoint_url: str, conn=None) -> MCPEndpointC
|
||||
mcp_client.set_websocket(websocket)
|
||||
|
||||
# 启动消息监听器
|
||||
mcp_client.listener_task = asyncio.create_task(
|
||||
_message_listener(mcp_client), name="xiaozhi-mcp-endpoint-listener"
|
||||
)
|
||||
asyncio.create_task(_message_listener(mcp_client))
|
||||
|
||||
# 发送初始化消息
|
||||
await send_mcp_endpoint_initialize(mcp_client)
|
||||
@@ -43,10 +39,6 @@ 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,25 +72,16 @@ class ServerMCPClient:
|
||||
|
||||
async def cleanup(self):
|
||||
"""清理MCP客户端资源"""
|
||||
task = self._worker_task
|
||||
if not task:
|
||||
if not self._worker_task:
|
||||
return
|
||||
|
||||
self._shutdown_evt.set()
|
||||
try:
|
||||
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:
|
||||
await asyncio.wait_for(self._worker_task, timeout=20)
|
||||
except (asyncio.TimeoutError, Exception) as e:
|
||||
self.logger.bind(tag=TAG).error(f"服务端MCP客户端关闭错误: {e}")
|
||||
finally:
|
||||
if task.done():
|
||||
self._worker_task = None
|
||||
self._worker_task = None
|
||||
|
||||
def has_tool(self, name: str) -> bool:
|
||||
"""检查是否包含指定工具
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
import os
|
||||
import json
|
||||
import re
|
||||
import uuid
|
||||
import queue
|
||||
@@ -17,6 +16,8 @@ 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,
|
||||
@@ -29,58 +30,6 @@ 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
|
||||
@@ -363,17 +312,13 @@ class TTSProviderBase(ABC):
|
||||
|
||||
# tts 消化线程
|
||||
self.tts_priority_thread = threading.Thread(
|
||||
target=self.tts_text_priority_thread,
|
||||
name=f"xiaozhi-tts-text-{id(self)}",
|
||||
daemon=True,
|
||||
target=self.tts_text_priority_thread, daemon=True
|
||||
)
|
||||
self.tts_priority_thread.start()
|
||||
|
||||
# 音频播放 消化线程
|
||||
self.audio_play_priority_thread = threading.Thread(
|
||||
target=self._audio_play_priority_thread,
|
||||
name=f"xiaozhi-tts-audio-{id(self)}",
|
||||
daemon=True,
|
||||
target=self._audio_play_priority_thread, daemon=True
|
||||
)
|
||||
self.audio_play_priority_thread.start()
|
||||
|
||||
@@ -424,8 +369,6 @@ 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
|
||||
@@ -474,15 +417,11 @@ 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 = getattr(
|
||||
self, "current_sentence_id", None
|
||||
) or getattr(self.conn, "sentence_id", None)
|
||||
sentence_id = None
|
||||
except queue.Empty:
|
||||
if self.conn.stop_event.is_set():
|
||||
break
|
||||
@@ -520,19 +459,7 @@ class TTSProviderBase(ABC):
|
||||
sendAudioMessage(self.conn, sentence_type, audio_datas, text, sentence_id),
|
||||
self.conn.loop,
|
||||
)
|
||||
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
|
||||
future.result()
|
||||
|
||||
# 记录输出和报告
|
||||
if self.conn.max_output_size > 0 and text:
|
||||
@@ -549,22 +476,6 @@ 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
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user