mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-26 00:53:54 +08:00
Merge branch 'main' into py-test
This commit is contained in:
@@ -177,5 +177,5 @@ public interface Constant {
|
|||||||
/**
|
/**
|
||||||
* 版本号
|
* 版本号
|
||||||
*/
|
*/
|
||||||
public static final String VERSION = "0.3.14";
|
public static final String VERSION = "0.4.2";
|
||||||
}
|
}
|
||||||
@@ -62,6 +62,13 @@ public class RedisKeys {
|
|||||||
return "agent:device:count:" + id;
|
return "agent:device:count:" + id;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* 获取智能体最后连接时间缓存key
|
||||||
|
*/
|
||||||
|
public static String getAgentDeviceLastConnectedAtById(String id) {
|
||||||
|
return "agent:device:lastConnected:" + id;
|
||||||
|
}
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* 获取系统配置缓存key
|
* 获取系统配置缓存key
|
||||||
*/
|
*/
|
||||||
|
|||||||
+6
@@ -1,12 +1,15 @@
|
|||||||
package xiaozhi.modules.agent.service.biz.impl;
|
package xiaozhi.modules.agent.service.biz.impl;
|
||||||
|
|
||||||
import java.util.Base64;
|
import java.util.Base64;
|
||||||
|
import java.util.Date;
|
||||||
|
|
||||||
import org.springframework.stereotype.Service;
|
import org.springframework.stereotype.Service;
|
||||||
import org.springframework.transaction.annotation.Transactional;
|
import org.springframework.transaction.annotation.Transactional;
|
||||||
|
|
||||||
import lombok.RequiredArgsConstructor;
|
import lombok.RequiredArgsConstructor;
|
||||||
import lombok.extern.slf4j.Slf4j;
|
import lombok.extern.slf4j.Slf4j;
|
||||||
|
import xiaozhi.common.redis.RedisKeys;
|
||||||
|
import xiaozhi.common.redis.RedisUtils;
|
||||||
import xiaozhi.modules.agent.dto.AgentChatHistoryReportDTO;
|
import xiaozhi.modules.agent.dto.AgentChatHistoryReportDTO;
|
||||||
import xiaozhi.modules.agent.entity.AgentChatHistoryEntity;
|
import xiaozhi.modules.agent.entity.AgentChatHistoryEntity;
|
||||||
import xiaozhi.modules.agent.entity.AgentEntity;
|
import xiaozhi.modules.agent.entity.AgentEntity;
|
||||||
@@ -29,6 +32,7 @@ public class AgentChatHistoryBizServiceImpl implements AgentChatHistoryBizServic
|
|||||||
private final AgentService agentService;
|
private final AgentService agentService;
|
||||||
private final AgentChatHistoryService agentChatHistoryService;
|
private final AgentChatHistoryService agentChatHistoryService;
|
||||||
private final AgentChatAudioService agentChatAudioService;
|
private final AgentChatAudioService agentChatAudioService;
|
||||||
|
private final RedisUtils redisUtils;
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* 处理聊天记录上报,包括文件上传和相关信息记录
|
* 处理聊天记录上报,包括文件上传和相关信息记录
|
||||||
@@ -77,6 +81,8 @@ public class AgentChatHistoryBizServiceImpl implements AgentChatHistoryBizServic
|
|||||||
|
|
||||||
// 3. 保存数据
|
// 3. 保存数据
|
||||||
agentChatHistoryService.save(entity);
|
agentChatHistoryService.save(entity);
|
||||||
|
// 4. 更新设备最后对话时间
|
||||||
|
redisUtils.set(RedisKeys.getAgentDeviceLastConnectedAtById(agentId), new Date());
|
||||||
return Boolean.TRUE;
|
return Boolean.TRUE;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+5
-2
@@ -21,6 +21,7 @@ import xiaozhi.modules.agent.dao.AgentDao;
|
|||||||
import xiaozhi.modules.agent.dto.AgentDTO;
|
import xiaozhi.modules.agent.dto.AgentDTO;
|
||||||
import xiaozhi.modules.agent.entity.AgentEntity;
|
import xiaozhi.modules.agent.entity.AgentEntity;
|
||||||
import xiaozhi.modules.agent.service.AgentService;
|
import xiaozhi.modules.agent.service.AgentService;
|
||||||
|
import xiaozhi.modules.device.service.DeviceService;
|
||||||
import xiaozhi.modules.model.service.ModelConfigService;
|
import xiaozhi.modules.model.service.ModelConfigService;
|
||||||
import xiaozhi.modules.security.user.SecurityUser;
|
import xiaozhi.modules.security.user.SecurityUser;
|
||||||
import xiaozhi.modules.sys.enums.SuperAdminEnum;
|
import xiaozhi.modules.sys.enums.SuperAdminEnum;
|
||||||
@@ -29,11 +30,11 @@ import xiaozhi.modules.timbre.service.TimbreService;
|
|||||||
@Service
|
@Service
|
||||||
@AllArgsConstructor
|
@AllArgsConstructor
|
||||||
public class AgentServiceImpl extends BaseServiceImpl<AgentDao, AgentEntity> implements AgentService {
|
public class AgentServiceImpl extends BaseServiceImpl<AgentDao, AgentEntity> implements AgentService {
|
||||||
|
|
||||||
private final AgentDao agentDao;
|
private final AgentDao agentDao;
|
||||||
private final TimbreService timbreModelService;
|
private final TimbreService timbreModelService;
|
||||||
private final ModelConfigService modelConfigService;
|
private final ModelConfigService modelConfigService;
|
||||||
private final RedisUtils redisUtils;
|
private final RedisUtils redisUtils;
|
||||||
|
private final DeviceService deviceService;
|
||||||
|
|
||||||
@Override
|
@Override
|
||||||
public PageData<AgentEntity> adminAgentList(Map<String, Object> params) {
|
public PageData<AgentEntity> adminAgentList(Map<String, Object> params) {
|
||||||
@@ -95,9 +96,11 @@ public class AgentServiceImpl extends BaseServiceImpl<AgentDao, AgentEntity> imp
|
|||||||
// 获取 TTS 音色名称
|
// 获取 TTS 音色名称
|
||||||
dto.setTtsVoiceName(timbreModelService.getTimbreNameById(agent.getTtsVoiceId()));
|
dto.setTtsVoiceName(timbreModelService.getTimbreNameById(agent.getTtsVoiceId()));
|
||||||
|
|
||||||
|
// 获取智能体最近的最后连接时长
|
||||||
|
dto.setLastConnectedAt(deviceService.getLatestLastConnectionTime(agent.getId()));
|
||||||
|
|
||||||
// 获取设备数量
|
// 获取设备数量
|
||||||
dto.setDeviceCount(getDeviceCountByAgentId(agent.getId()));
|
dto.setDeviceCount(getDeviceCountByAgentId(agent.getId()));
|
||||||
|
|
||||||
return dto;
|
return dto;
|
||||||
}).collect(Collectors.toList());
|
}).collect(Collectors.toList());
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -6,6 +6,16 @@ import com.baomidou.mybatisplus.core.mapper.BaseMapper;
|
|||||||
|
|
||||||
import xiaozhi.modules.device.entity.DeviceEntity;
|
import xiaozhi.modules.device.entity.DeviceEntity;
|
||||||
|
|
||||||
|
import java.util.Date;
|
||||||
|
import java.util.List;
|
||||||
|
|
||||||
@Mapper
|
@Mapper
|
||||||
public interface DeviceDao extends BaseMapper<DeviceEntity> {
|
public interface DeviceDao extends BaseMapper<DeviceEntity> {
|
||||||
|
/**
|
||||||
|
* 获取此智能体全部设备的最后连接时间
|
||||||
|
* @param agentId 智能体id
|
||||||
|
* @return
|
||||||
|
*/
|
||||||
|
Date getAllLastConnectedAtByAgentId(String agentId);
|
||||||
|
|
||||||
}
|
}
|
||||||
@@ -1,5 +1,6 @@
|
|||||||
package xiaozhi.modules.device.service;
|
package xiaozhi.modules.device.service;
|
||||||
|
|
||||||
|
import java.util.Date;
|
||||||
import java.util.List;
|
import java.util.List;
|
||||||
|
|
||||||
import xiaozhi.common.page.PageData;
|
import xiaozhi.common.page.PageData;
|
||||||
@@ -78,4 +79,13 @@ public interface DeviceService extends BaseService<DeviceEntity> {
|
|||||||
* @return 激活码
|
* @return 激活码
|
||||||
*/
|
*/
|
||||||
String geCodeByDeviceId(String deviceId);
|
String geCodeByDeviceId(String deviceId);
|
||||||
|
|
||||||
|
/**
|
||||||
|
* 获取这个智能体设备理的最近的最后连接时间
|
||||||
|
* @param agentId 智能体id
|
||||||
|
* @return 返回设备最近的最后连接时间
|
||||||
|
*/
|
||||||
|
Date getLatestLastConnectionTime(String agentId);
|
||||||
|
|
||||||
|
|
||||||
}
|
}
|
||||||
+14
@@ -258,6 +258,20 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
|
|||||||
return null;
|
return null;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
@Override
|
||||||
|
public Date getLatestLastConnectionTime(String agentId) {
|
||||||
|
// 查询是否有缓存时间,有则返回
|
||||||
|
Date cachedDate = (Date) redisUtils.get(RedisKeys.getAgentDeviceLastConnectedAtById(agentId));
|
||||||
|
if (cachedDate != null) {
|
||||||
|
return cachedDate;
|
||||||
|
}
|
||||||
|
Date maxDate = deviceDao.getAllLastConnectedAtByAgentId(agentId);
|
||||||
|
if (maxDate != null) {
|
||||||
|
redisUtils.set(RedisKeys.getAgentDeviceLastConnectedAtById(agentId), maxDate);
|
||||||
|
}
|
||||||
|
return maxDate;
|
||||||
|
}
|
||||||
|
|
||||||
private String getDeviceCacheKey(String deviceId) {
|
private String getDeviceCacheKey(String deviceId) {
|
||||||
String safeDeviceId = deviceId.replace(":", "_").toLowerCase();
|
String safeDeviceId = deviceId.replace(":", "_").toLowerCase();
|
||||||
String dataKey = String.format("ota:activation:data:%s", safeDeviceId);
|
String dataKey = String.format("ota:activation:data:%s", safeDeviceId);
|
||||||
|
|||||||
@@ -0,0 +1,43 @@
|
|||||||
|
-- 添加百度ASR模型配置
|
||||||
|
delete from `ai_model_config` where `id` = 'ASR_BaiduASR';
|
||||||
|
INSERT INTO `ai_model_config` VALUES ('ASR_BaiduASR', 'ASR', 'BaiduASR', '百度语音识别', 0, 1, '{\"type\": \"baidu\", \"app_id\": \"\", \"api_key\": \"\", \"secret_key\": \"\", \"dev_pid\": 1537, \"output_dir\": \"tmp/\"}', NULL, NULL, 7, NULL, NULL, NULL, NULL);
|
||||||
|
|
||||||
|
|
||||||
|
-- 添加百度ASR供应器
|
||||||
|
delete from `ai_model_provider` where `id` = 'SYSTEM_ASR_BaiduASR';
|
||||||
|
INSERT INTO `ai_model_provider` (`id`, `model_type`, `provider_code`, `name`, `fields`, `sort`, `creator`, `create_date`, `updater`, `update_date`) VALUES
|
||||||
|
('SYSTEM_ASR_BaiduASR', 'ASR', 'baidu', '百度语音识别', '[{"key":"app_id","label":"应用AppID","type":"string"},{"key":"api_key","label":"API Key","type":"string"},{"key":"secret_key","label":"Secret Key","type":"string"},{"key":"dev_pid","label":"语言参数","type":"number"},{"key":"output_dir","label":"输出目录","type":"string"}]', 7, 1, NOW(), 1, NOW());
|
||||||
|
|
||||||
|
|
||||||
|
-- 更新百度ASR配置说明
|
||||||
|
UPDATE `ai_model_config` SET
|
||||||
|
`doc_link` = 'https://console.bce.baidu.com/ai-engine/old/#/ai/speech/app/list',
|
||||||
|
`remark` = '百度ASR配置说明:
|
||||||
|
1. 访问 https://console.bce.baidu.com/ai-engine/old/#/ai/speech/app/list
|
||||||
|
2. 创建新应用
|
||||||
|
3. 获取AppID、API Key和Secret Key
|
||||||
|
4. 填入配置文件中
|
||||||
|
查看资源额度:https://console.bce.baidu.com/ai-engine/old/#/ai/speech/overview/resource/list
|
||||||
|
语言参数说明:https://ai.baidu.com/ai-doc/SPEECH/0lbxfnc9b
|
||||||
|
' WHERE `id` = 'ASR_BaiduASR';
|
||||||
|
|
||||||
|
-- 更新豆包供应器字段
|
||||||
|
update `ai_model_provider` set `fields` =
|
||||||
|
'[{"key":"appid","label":"应用ID","type":"string"},{"key":"access_token","label":"访问令牌","type":"string"},{"key":"cluster","label":"集群","type":"string"},{"key":"boosting_table_name","label":"热词文件名称","type":"string"},{"key":"correct_table_name","label":"替换词文件名称","type":"string"},{"key":"output_dir","label":"输出目录","type":"string"}]'
|
||||||
|
where `id` = 'SYSTEM_ASR_DoubaoASR';
|
||||||
|
|
||||||
|
-- 更新豆包ASR配置说明
|
||||||
|
UPDATE `ai_model_config` SET
|
||||||
|
`doc_link` = 'https://console.volcengine.com/speech/app',
|
||||||
|
`remark` = '豆包ASR配置说明:
|
||||||
|
1. 需要在火山引擎控制台创建应用并获取appid和access_token
|
||||||
|
2. 支持中文语音识别
|
||||||
|
3. 需要网络连接
|
||||||
|
4. 输出文件保存在tmp/目录
|
||||||
|
申请步骤:
|
||||||
|
1. 访问 https://console.volcengine.com/speech/app
|
||||||
|
2. 创建新应用
|
||||||
|
3. 获取appid和access_token
|
||||||
|
4. 填入配置文件中
|
||||||
|
如需设置热词,请参考:https://www.volcengine.com/docs/6561/155738
|
||||||
|
' WHERE `id` = 'ASR_DoubaoASR';
|
||||||
@@ -99,4 +99,11 @@ databaseChangeLog:
|
|||||||
changes:
|
changes:
|
||||||
- sqlFile:
|
- sqlFile:
|
||||||
encoding: utf8
|
encoding: utf8
|
||||||
path: classpath:db/changelog/202505022134.sql
|
path: classpath:db/changelog/202505022134.sql
|
||||||
|
- changeSet:
|
||||||
|
id: 202505081146
|
||||||
|
author: hrz
|
||||||
|
changes:
|
||||||
|
- sqlFile:
|
||||||
|
encoding: utf8
|
||||||
|
path: classpath:db/changelog/202505081146.sql
|
||||||
@@ -0,0 +1,12 @@
|
|||||||
|
<?xml version="1.0" encoding="UTF-8"?>
|
||||||
|
<!DOCTYPE mapper PUBLIC "-//mybatis.org//DTD Mapper 3.0//EN" "http://mybatis.org/dtd/mybatis-3-mapper.dtd">
|
||||||
|
<mapper namespace="xiaozhi.modules.device.dao.DeviceDao">
|
||||||
|
<!-- 获取此智能体全部设备的最后连接时间 -->
|
||||||
|
<select id="getAllLastConnectedAtByAgentId" resultType="java.util.Date">
|
||||||
|
SELECT last_connected_at FROM ai_device
|
||||||
|
WHERE
|
||||||
|
agent_id = #{agentId}
|
||||||
|
order by
|
||||||
|
last_connected_at desc limit 0,1
|
||||||
|
</select>
|
||||||
|
</mapper>
|
||||||
@@ -1,5 +1,8 @@
|
|||||||
package xiaozhi.modules.device;
|
package xiaozhi.modules.device;
|
||||||
|
|
||||||
|
import java.util.HashMap;
|
||||||
|
import java.util.UUID;
|
||||||
|
|
||||||
import org.junit.jupiter.api.Assertions;
|
import org.junit.jupiter.api.Assertions;
|
||||||
import org.junit.jupiter.api.DisplayName;
|
import org.junit.jupiter.api.DisplayName;
|
||||||
import org.junit.jupiter.api.Test;
|
import org.junit.jupiter.api.Test;
|
||||||
@@ -8,8 +11,9 @@ import org.springframework.boot.test.context.SpringBootTest;
|
|||||||
import org.springframework.test.context.ActiveProfiles;
|
import org.springframework.test.context.ActiveProfiles;
|
||||||
|
|
||||||
import lombok.extern.slf4j.Slf4j;
|
import lombok.extern.slf4j.Slf4j;
|
||||||
import xiaozhi.common.redis.RedisKeys;
|
|
||||||
import xiaozhi.common.redis.RedisUtils;
|
import xiaozhi.common.redis.RedisUtils;
|
||||||
|
import xiaozhi.modules.sys.dto.SysUserDTO;
|
||||||
|
import xiaozhi.modules.sys.service.SysUserService;
|
||||||
|
|
||||||
@Slf4j
|
@Slf4j
|
||||||
@SpringBootTest
|
@SpringBootTest
|
||||||
@@ -19,18 +23,37 @@ public class DeviceTest {
|
|||||||
|
|
||||||
@Autowired
|
@Autowired
|
||||||
private RedisUtils redisUtils;
|
private RedisUtils redisUtils;
|
||||||
|
@Autowired
|
||||||
|
private SysUserService sysUserService;
|
||||||
|
|
||||||
|
@Test
|
||||||
|
public void testSaveUser() {
|
||||||
|
SysUserDTO userDTO = new SysUserDTO();
|
||||||
|
userDTO.setUsername("test");
|
||||||
|
userDTO.setPassword(UUID.randomUUID().toString());
|
||||||
|
sysUserService.save(userDTO);
|
||||||
|
}
|
||||||
|
|
||||||
@Test
|
@Test
|
||||||
@DisplayName("测试写入设备信息")
|
@DisplayName("测试写入设备信息")
|
||||||
public void testWriteDeviceInfo() {
|
public void testWriteDeviceInfo() {
|
||||||
log.info("开始测试写入设备信息...");
|
log.info("开始测试写入设备信息...");
|
||||||
|
|
||||||
// 模拟设备MAC地址
|
// 模拟设备MAC地址
|
||||||
String macAddress = "00:11:22:33:44:55";
|
String macAddress = "00:11:22:33:44:66";
|
||||||
// 模拟设备验证码
|
// 模拟设备验证码
|
||||||
String deviceCode = "123456";
|
String deviceCode = "123456";
|
||||||
|
|
||||||
String redisKey = RedisKeys.getDeviceCaptchaKey(deviceCode);
|
HashMap<String, Object> map = new HashMap<>();
|
||||||
|
map.put("mac_address", macAddress);
|
||||||
|
map.put("activation_code", deviceCode);
|
||||||
|
map.put("board", "硬件型号");
|
||||||
|
map.put("app_version", "0.3.13");
|
||||||
|
|
||||||
|
String safeDeviceId = macAddress.replace(":", "_").toLowerCase();
|
||||||
|
String cacheDeviceKey = String.format("ota:activation:data:%s", safeDeviceId);
|
||||||
|
redisUtils.set(cacheDeviceKey, map, 300);
|
||||||
|
|
||||||
|
String redisKey = "ota:activation:code:" + deviceCode;
|
||||||
log.info("Redis Key: {}", redisKey);
|
log.info("Redis Key: {}", redisKey);
|
||||||
|
|
||||||
// 将设备信息写入Redis
|
// 将设备信息写入Redis
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
<template>
|
<template>
|
||||||
<el-dialog :visible="visible" @close="handleClose" width="400px" center>
|
<el-dialog :visible="visible" @close="handleClose" width="24%" center>
|
||||||
<div
|
<div
|
||||||
style="margin: 0 10px 10px;display: flex;align-items: center;gap: 10px;font-weight: 700;font-size: 20px;text-align: left;color: #3d4566;">
|
style="margin: 0 10px 10px;display: flex;align-items: center;gap: 10px;font-weight: 700;font-size: 20px;text-align: left;color: #3d4566;">
|
||||||
<div
|
<div
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
<template>
|
<template>
|
||||||
<el-dialog :visible="dialogVisible" @update:visible="handleVisibleChange" width="975px" center
|
<el-dialog :visible="dialogVisible" @update:visible="handleVisibleChange" width="57%" center
|
||||||
custom-class="custom-dialog" :show-close="false" class="center-dialog">
|
custom-class="custom-dialog" :show-close="false" class="center-dialog">
|
||||||
<div style="margin: 0 18px; text-align: left; padding: 10px; border-radius: 10px;">
|
<div style="margin: 0 18px; text-align: left; padding: 10px; border-radius: 10px;">
|
||||||
<div style="font-size: 30px; color: #3d4566; margin-top: -10px; margin-bottom: 10px; text-align: center;">
|
<div style="font-size: 30px; color: #3d4566; margin-top: -10px; margin-bottom: 10px; text-align: center;">
|
||||||
@@ -54,7 +54,7 @@
|
|||||||
</el-form-item>
|
</el-form-item>
|
||||||
|
|
||||||
<el-form-item label="备注" prop="remark" class="prop-remark">
|
<el-form-item label="备注" prop="remark" class="prop-remark">
|
||||||
<el-input v-model="formData.remark" type="textarea" :rows="3" placeholder="请输入模型备注"
|
<el-input v-model="formData.remark" type="textarea" :rows="3" placeholder="请输入模型备注" :autosize="{ minRows: 3, maxRows: 5 }"
|
||||||
class="custom-input-bg"></el-input>
|
class="custom-input-bg"></el-input>
|
||||||
</el-form-item>
|
</el-form-item>
|
||||||
</el-form>
|
</el-form>
|
||||||
@@ -271,7 +271,7 @@ export default {
|
|||||||
}
|
}
|
||||||
|
|
||||||
.center-dialog .el-dialog {
|
.center-dialog .el-dialog {
|
||||||
margin: 4% 0 auto !important;
|
margin: 0 0 auto !important;
|
||||||
display: flex;
|
display: flex;
|
||||||
flex-direction: column;
|
flex-direction: column;
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
<template>
|
<template>
|
||||||
<el-dialog :visible="visible" @close="handleClose" width="400px" center @open="handleOpen">
|
<el-dialog :visible="visible" @close="handleClose" width="25%" center @open="handleOpen">
|
||||||
<div
|
<div
|
||||||
style="margin: 0 10px 10px;display: flex;align-items: center;gap: 10px;font-weight: 700;font-size: 20px;text-align: left;color: #3d4566;">
|
style="margin: 0 10px 10px;display: flex;align-items: center;gap: 10px;font-weight: 700;font-size: 20px;text-align: left;color: #3d4566;">
|
||||||
<div
|
<div
|
||||||
@@ -10,7 +10,7 @@
|
|||||||
</div>
|
</div>
|
||||||
<div style="height: 1px;background: #e8f0ff;" />
|
<div style="height: 1px;background: #e8f0ff;" />
|
||||||
<div style="margin: 22px 15px;">
|
<div style="margin: 22px 15px;">
|
||||||
<div style="font-weight: 400;font-size: 14px;text-align: left;color: #3d4566;">
|
<div style="font-weight: 400;text-align: left;color: #3d4566;">
|
||||||
<div style="color: red;display: inline-block;">*</div> 智能体名称:
|
<div style="color: red;display: inline-block;">*</div> 智能体名称:
|
||||||
</div>
|
</div>
|
||||||
<div class="input-46" style="margin-top: 12px;">
|
<div class="input-46" style="margin-top: 12px;">
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
<template>
|
<template>
|
||||||
<form>
|
<form>
|
||||||
<el-dialog :visible.sync="value" width="400px" center>
|
<el-dialog :visible.sync="dialogVisible" width="24%" center>
|
||||||
<div
|
<div
|
||||||
style="margin: 0 10px 10px;display: flex;align-items: center;gap: 10px;font-weight: 700;font-size: 20px;text-align: left;color: #3d4566;">
|
style="margin: 0 10px 10px;display: flex;align-items: center;gap: 10px;font-weight: 700;font-size: 20px;text-align: left;color: #3d4566;">
|
||||||
<div
|
<div
|
||||||
@@ -60,11 +60,20 @@ export default {
|
|||||||
},
|
},
|
||||||
data() {
|
data() {
|
||||||
return {
|
return {
|
||||||
|
dialogVisible: this.value,
|
||||||
oldPassword: "",
|
oldPassword: "",
|
||||||
newPassword: "",
|
newPassword: "",
|
||||||
confirmNewPassword: ""
|
confirmNewPassword: ""
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
|
watch: {
|
||||||
|
value(val) {
|
||||||
|
this.dialogVisible = val;
|
||||||
|
},
|
||||||
|
dialogVisible(val) {
|
||||||
|
this.$emit('input', val);
|
||||||
|
}
|
||||||
|
},
|
||||||
methods: {
|
methods: {
|
||||||
...mapActions(['logout']), // 引入Vuex的logout action
|
...mapActions(['logout']), // 引入Vuex的logout action
|
||||||
confirm() {
|
confirm() {
|
||||||
@@ -101,7 +110,7 @@ export default {
|
|||||||
this.$emit('input', false);
|
this.$emit('input', false);
|
||||||
},
|
},
|
||||||
cancel() {
|
cancel() {
|
||||||
this.$emit('input', false);
|
this.dialogVisible = false;
|
||||||
this.resetForm();
|
this.resetForm();
|
||||||
},
|
},
|
||||||
resetForm() {
|
resetForm() {
|
||||||
|
|||||||
@@ -31,7 +31,7 @@
|
|||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
<div class="version-info">
|
<div class="version-info">
|
||||||
<div>最近对话:{{ device.lastConnectedAt }}</div>
|
<div>最近对话:{{ formattedLastConnectedTime }}</div>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
</template>
|
</template>
|
||||||
@@ -45,6 +45,27 @@ export default {
|
|||||||
data() {
|
data() {
|
||||||
return { switchValue: false }
|
return { switchValue: false }
|
||||||
},
|
},
|
||||||
|
computed: {
|
||||||
|
formattedLastConnectedTime() {
|
||||||
|
if (!this.device.lastConnectedAt) return '暂未对话';
|
||||||
|
|
||||||
|
const lastTime = new Date(this.device.lastConnectedAt);
|
||||||
|
const now = new Date();
|
||||||
|
const diffMinutes = Math.floor((now - lastTime) / (1000 * 60));
|
||||||
|
|
||||||
|
if (diffMinutes <= 1) {
|
||||||
|
return '刚刚';
|
||||||
|
} else if (diffMinutes < 60) {
|
||||||
|
return `${diffMinutes}分钟前`;
|
||||||
|
} else if (diffMinutes < 24 * 60) {
|
||||||
|
const hours = Math.floor(diffMinutes / 60);
|
||||||
|
const minutes = diffMinutes % 60;
|
||||||
|
return `${hours}小时${minutes > 0 ? minutes + '分钟' : ''}前`;
|
||||||
|
} else {
|
||||||
|
return this.device.lastConnectedAt;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
},
|
||||||
methods: {
|
methods: {
|
||||||
handleDelete() {
|
handleDelete() {
|
||||||
this.$emit('delete', this.device.agentId)
|
this.$emit('delete', this.device.agentId)
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
<template>
|
<template>
|
||||||
<el-dialog :title="title" :visible.sync="visible" width="500px" @close="handleClose">
|
<el-dialog :title="title" :visible.sync="dialogVisible" width="30%" @close="handleClose">
|
||||||
<el-form :model="form" :rules="rules" ref="form" label-width="100px">
|
<el-form :model="form" :rules="rules" ref="form" label-width="100px">
|
||||||
<el-form-item label="字典标签" prop="dictLabel">
|
<el-form-item label="字典标签" prop="dictLabel">
|
||||||
<el-input v-model="form.dictLabel" placeholder="请输入字典标签"></el-input>
|
<el-input v-model="form.dictLabel" placeholder="请输入字典标签"></el-input>
|
||||||
@@ -41,6 +41,7 @@ export default {
|
|||||||
},
|
},
|
||||||
data() {
|
data() {
|
||||||
return {
|
return {
|
||||||
|
dialogVisible: this.visible,
|
||||||
form: {
|
form: {
|
||||||
id: null,
|
id: null,
|
||||||
dictTypeId: null,
|
dictTypeId: null,
|
||||||
@@ -70,12 +71,18 @@ export default {
|
|||||||
}
|
}
|
||||||
},
|
},
|
||||||
immediate: true
|
immediate: true
|
||||||
|
},
|
||||||
|
visible(val) {
|
||||||
|
this.dialogVisible = val;
|
||||||
|
},
|
||||||
|
dialogVisible(val) {
|
||||||
|
this.$emit('update:visible', val);
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
methods: {
|
methods: {
|
||||||
handleClose() {
|
handleClose() {
|
||||||
this.$emit('update:visible', false)
|
this.dialogVisible = false;
|
||||||
this.resetForm()
|
this.resetForm();
|
||||||
},
|
},
|
||||||
resetForm() {
|
resetForm() {
|
||||||
this.form = {
|
this.form = {
|
||||||
@@ -102,4 +109,8 @@ export default {
|
|||||||
.dialog-footer {
|
.dialog-footer {
|
||||||
text-align: right;
|
text-align: right;
|
||||||
}
|
}
|
||||||
|
:deep(.el-dialog) {
|
||||||
|
border-radius: 15px;
|
||||||
|
}
|
||||||
|
|
||||||
</style>
|
</style>
|
||||||
@@ -1,5 +1,5 @@
|
|||||||
<template>
|
<template>
|
||||||
<el-dialog :title="title" :visible.sync="visible" width="500px" @close="handleClose">
|
<el-dialog :title="title" :visible.sync="dialogVisible" width="30%" @close="handleClose">
|
||||||
<el-form :model="form" :rules="rules" ref="form" label-width="120px">
|
<el-form :model="form" :rules="rules" ref="form" label-width="120px">
|
||||||
<el-form-item label="字典类型名称" prop="dictName">
|
<el-form-item label="字典类型名称" prop="dictName">
|
||||||
<el-input v-model="form.dictName" placeholder="请输入字典类型名称"></el-input>
|
<el-input v-model="form.dictName" placeholder="请输入字典类型名称"></el-input>
|
||||||
@@ -34,6 +34,7 @@ export default {
|
|||||||
},
|
},
|
||||||
data() {
|
data() {
|
||||||
return {
|
return {
|
||||||
|
dialogVisible: this.visible,
|
||||||
form: {
|
form: {
|
||||||
id: null,
|
id: null,
|
||||||
dictName: '',
|
dictName: '',
|
||||||
@@ -46,6 +47,12 @@ export default {
|
|||||||
}
|
}
|
||||||
},
|
},
|
||||||
watch: {
|
watch: {
|
||||||
|
visible(val) {
|
||||||
|
this.dialogVisible = val;
|
||||||
|
},
|
||||||
|
dialogVisible(val) {
|
||||||
|
this.$emit('update:visible', val);
|
||||||
|
},
|
||||||
dictTypeData: {
|
dictTypeData: {
|
||||||
handler(val) {
|
handler(val) {
|
||||||
if (val) {
|
if (val) {
|
||||||
@@ -57,7 +64,7 @@ export default {
|
|||||||
},
|
},
|
||||||
methods: {
|
methods: {
|
||||||
handleClose() {
|
handleClose() {
|
||||||
this.$emit('update:visible', false)
|
this.dialogVisible = false;
|
||||||
this.resetForm()
|
this.resetForm()
|
||||||
},
|
},
|
||||||
resetForm() {
|
resetForm() {
|
||||||
@@ -83,4 +90,8 @@ export default {
|
|||||||
.dialog-footer {
|
.dialog-footer {
|
||||||
text-align: right;
|
text-align: right;
|
||||||
}
|
}
|
||||||
|
:deep(.el-dialog) {
|
||||||
|
border-radius: 15px;
|
||||||
|
}
|
||||||
|
|
||||||
</style>
|
</style>
|
||||||
@@ -1,5 +1,5 @@
|
|||||||
<template>
|
<template>
|
||||||
<el-dialog :title="title" :visible.sync="visible" width="500px" @close="handleClose" @open="handleOpen">
|
<el-dialog :title="title" :visible.sync="dialogVisible" width="30%" @close="handleClose" @open="handleOpen">
|
||||||
<el-form ref="form" :model="form" :rules="rules" label-width="100px">
|
<el-form ref="form" :model="form" :rules="rules" label-width="100px">
|
||||||
<el-form-item label="固件名称" prop="firmwareName">
|
<el-form-item label="固件名称" prop="firmwareName">
|
||||||
<el-input v-model="form.firmwareName" placeholder="请输入固件名称(板子+版本号)"></el-input>
|
<el-input v-model="form.firmwareName" placeholder="请输入固件名称(板子+版本号)"></el-input>
|
||||||
@@ -59,11 +59,13 @@ export default {
|
|||||||
default: () => []
|
default: () => []
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
|
|
||||||
data() {
|
data() {
|
||||||
return {
|
return {
|
||||||
uploadProgress: 0,
|
uploadProgress: 0,
|
||||||
uploadStatus: '',
|
uploadStatus: '',
|
||||||
isUploading: false,
|
isUploading: false,
|
||||||
|
dialogVisible: this.visible,
|
||||||
rules: {
|
rules: {
|
||||||
firmwareName: [
|
firmwareName: [
|
||||||
{ required: true, message: '请输入固件名称(板子+版本号)', trigger: 'blur' }
|
{ required: true, message: '请输入固件名称(板子+版本号)', trigger: 'blur' }
|
||||||
@@ -90,10 +92,18 @@ export default {
|
|||||||
created() {
|
created() {
|
||||||
// 移除 getDictDataByType 调用
|
// 移除 getDictDataByType 调用
|
||||||
},
|
},
|
||||||
|
watch: {
|
||||||
|
visible(val) {
|
||||||
|
this.dialogVisible = val;
|
||||||
|
},
|
||||||
|
dialogVisible(val) {
|
||||||
|
this.$emit('update:visible', val);
|
||||||
|
},
|
||||||
|
},
|
||||||
methods: {
|
methods: {
|
||||||
// 移除 getFirmwareTypes 方法
|
// 移除 getFirmwareTypes 方法
|
||||||
handleClose() {
|
handleClose() {
|
||||||
this.$refs.form.clearValidate();
|
this.dialogVisible = false;
|
||||||
this.$emit('cancel');
|
this.$emit('cancel');
|
||||||
},
|
},
|
||||||
handleCancel() {
|
handleCancel() {
|
||||||
@@ -201,13 +211,17 @@ export default {
|
|||||||
</script>
|
</script>
|
||||||
|
|
||||||
<style lang="scss" scoped>
|
<style lang="scss" scoped>
|
||||||
|
::v-deep .el-dialog {
|
||||||
|
border-radius: 20px;
|
||||||
|
}
|
||||||
|
|
||||||
.upload-demo {
|
.upload-demo {
|
||||||
text-align: left;
|
text-align: left;
|
||||||
}
|
}
|
||||||
|
|
||||||
.el-upload__tip {
|
.el-upload__tip {
|
||||||
line-height: 1.2;
|
line-height: 1.2;
|
||||||
padding-top: 5px;
|
padding-top: 2%;
|
||||||
color: #909399;
|
color: #909399;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
<template>
|
<template>
|
||||||
<el-dialog :visible.sync="dialogVisible" width="975px" center custom-class="custom-dialog" :show-close="false"
|
<el-dialog :visible.sync="dialogVisible" width="57%" center custom-class="custom-dialog" :show-close="false"
|
||||||
class="center-dialog">
|
class="center-dialog" >
|
||||||
<div style="margin: 0 18px; text-align: left; padding: 10px; border-radius: 10px;">
|
<div style="margin: 0 18px; text-align: left; padding: 10px; border-radius: 10px;">
|
||||||
<div style="font-size: 30px; color: #3d4566; margin-top: -10px; margin-bottom: 10px; text-align: center;">
|
<div style="font-size: 30px; color: #3d4566; margin-top: -10px; margin-bottom: 10px; text-align: center;">
|
||||||
修改模型
|
修改模型
|
||||||
@@ -53,7 +53,7 @@
|
|||||||
</el-form-item>
|
</el-form-item>
|
||||||
|
|
||||||
<el-form-item label="备注" prop="remark" class="prop-remark">
|
<el-form-item label="备注" prop="remark" class="prop-remark">
|
||||||
<el-input v-model="form.remark" type="textarea" :rows="3" placeholder="请输入模型备注"
|
<el-input v-model="form.remark" type="textarea" :rows="3" placeholder="请输入模型备注" :autosize="{ minRows: 3, maxRows: 5 }"
|
||||||
class="custom-input-bg"></el-input>
|
class="custom-input-bg"></el-input>
|
||||||
</el-form-item>
|
</el-form-item>
|
||||||
</el-form>
|
</el-form>
|
||||||
@@ -296,7 +296,7 @@ export default {
|
|||||||
};
|
};
|
||||||
</script>
|
</script>
|
||||||
|
|
||||||
<style scoped>
|
<style lang="scss" scoped>
|
||||||
.custom-dialog {
|
.custom-dialog {
|
||||||
position: relative;
|
position: relative;
|
||||||
border-radius: 20px;
|
border-radius: 20px;
|
||||||
@@ -316,11 +316,6 @@ export default {
|
|||||||
justify-content: center;
|
justify-content: center;
|
||||||
}
|
}
|
||||||
|
|
||||||
.center-dialog .el-dialog {
|
|
||||||
margin: 4% 0 auto !important;
|
|
||||||
display: flex;
|
|
||||||
flex-direction: column;
|
|
||||||
}
|
|
||||||
|
|
||||||
.custom-close-btn {
|
.custom-close-btn {
|
||||||
position: absolute;
|
position: absolute;
|
||||||
|
|||||||
@@ -487,7 +487,7 @@ export default {
|
|||||||
};
|
};
|
||||||
</script>
|
</script>
|
||||||
|
|
||||||
<style scoped>
|
<style lang="scss" scoped>
|
||||||
|
|
||||||
::v-deep .el-dialog {
|
::v-deep .el-dialog {
|
||||||
border-radius: 8px !important;
|
border-radius: 8px !important;
|
||||||
@@ -648,12 +648,6 @@ export default {
|
|||||||
margin: 0 auto;
|
margin: 0 auto;
|
||||||
}
|
}
|
||||||
|
|
||||||
/* 新增按钮组样式 */
|
|
||||||
.action-buttons {
|
|
||||||
bottom: 20px;
|
|
||||||
padding-top: 10px;
|
|
||||||
}
|
|
||||||
|
|
||||||
.action-buttons .el-button {
|
.action-buttons .el-button {
|
||||||
padding: 8px 15px;
|
padding: 8px 15px;
|
||||||
font-size: 11px;
|
font-size: 11px;
|
||||||
@@ -692,7 +686,6 @@ export default {
|
|||||||
position: static;
|
position: static;
|
||||||
padding: 15px 0;
|
padding: 15px 0;
|
||||||
background: white;
|
background: white;
|
||||||
box-shadow: 0 -2px 12px rgba(0,0,0,0.05);
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/* 输入框自适应 */
|
/* 输入框自适应 */
|
||||||
|
|||||||
@@ -49,9 +49,13 @@
|
|||||||
v-loading="dictDataLoading" element-loading-text="拼命加载中"
|
v-loading="dictDataLoading" element-loading-text="拼命加载中"
|
||||||
element-loading-spinner="el-icon-loading"
|
element-loading-spinner="el-icon-loading"
|
||||||
element-loading-background="rgba(255, 255, 255, 0.7)"
|
element-loading-background="rgba(255, 255, 255, 0.7)"
|
||||||
@selection-change="handleDictDataSelectionChange" class="data-table"
|
class="data-table"
|
||||||
header-row-class-name="table-header">
|
header-row-class-name="table-header">
|
||||||
<el-table-column type="selection" width="55" align="center"></el-table-column>
|
<el-table-column label="选择" align="center" width="55">
|
||||||
|
<template slot-scope="scope">
|
||||||
|
<el-checkbox v-model="scope.row.selected"></el-checkbox>
|
||||||
|
</template>
|
||||||
|
</el-table-column>
|
||||||
<el-table-column label="字典标签" prop="dictLabel" align="center"></el-table-column>
|
<el-table-column label="字典标签" prop="dictLabel" align="center"></el-table-column>
|
||||||
<el-table-column label="字典值" prop="dictValue" align="center"></el-table-column>
|
<el-table-column label="字典值" prop="dictValue" align="center"></el-table-column>
|
||||||
<el-table-column label="排序" prop="sort" align="center"></el-table-column>
|
<el-table-column label="排序" prop="sort" align="center"></el-table-column>
|
||||||
@@ -153,7 +157,6 @@ export default {
|
|||||||
// 字典数据相关
|
// 字典数据相关
|
||||||
dictDataList: [],
|
dictDataList: [],
|
||||||
dictDataLoading: false,
|
dictDataLoading: false,
|
||||||
selectedDictData: [],
|
|
||||||
isAllDictDataSelected: false,
|
isAllDictDataSelected: false,
|
||||||
dictDataDialogVisible: false,
|
dictDataDialogVisible: false,
|
||||||
dictDataDialogTitle: '新增字典数据',
|
dictDataDialogTitle: '新增字典数据',
|
||||||
@@ -265,7 +268,10 @@ export default {
|
|||||||
dictValue: ''
|
dictValue: ''
|
||||||
}, ({ data }) => {
|
}, ({ data }) => {
|
||||||
if (data.code === 0) {
|
if (data.code === 0) {
|
||||||
this.dictDataList = data.data.list
|
this.dictDataList = data.data.list.map(item => ({
|
||||||
|
...item,
|
||||||
|
selected: false
|
||||||
|
}))
|
||||||
this.total = data.data.total
|
this.total = data.data.total
|
||||||
} else {
|
} else {
|
||||||
this.$message.error(data.msg || '获取字典数据失败')
|
this.$message.error(data.msg || '获取字典数据失败')
|
||||||
@@ -273,16 +279,11 @@ export default {
|
|||||||
this.dictDataLoading = false
|
this.dictDataLoading = false
|
||||||
})
|
})
|
||||||
},
|
},
|
||||||
handleDictDataSelectionChange(val) {
|
|
||||||
this.selectedDictData = val
|
|
||||||
this.isAllDictDataSelected = val.length === this.dictDataList.length
|
|
||||||
},
|
|
||||||
selectAllDictData() {
|
selectAllDictData() {
|
||||||
if (this.isAllDictDataSelected) {
|
this.isAllDictDataSelected = !this.isAllDictDataSelected
|
||||||
this.$refs.dictDataTable.clearSelection()
|
this.dictDataList.forEach(row => {
|
||||||
} else {
|
row.selected = this.isAllDictDataSelected
|
||||||
this.$refs.dictDataTable.toggleAllSelection()
|
})
|
||||||
}
|
|
||||||
},
|
},
|
||||||
showAddDictDataDialog() {
|
showAddDictDataDialog() {
|
||||||
if (!this.selectedDictType) {
|
if (!this.selectedDictType) {
|
||||||
@@ -329,17 +330,18 @@ export default {
|
|||||||
})
|
})
|
||||||
},
|
},
|
||||||
batchDeleteDictData() {
|
batchDeleteDictData() {
|
||||||
if (this.selectedDictData.length === 0) {
|
const selectedRows = this.dictDataList.filter(row => row.selected)
|
||||||
|
if (selectedRows.length === 0) {
|
||||||
this.$message.warning('请选择要删除的字典数据')
|
this.$message.warning('请选择要删除的字典数据')
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
this.$confirm('确定要删除选中的字典数据吗?', '提示', {
|
this.$confirm(`确定要删除选中的${selectedRows.length}个字典数据吗?`, '提示', {
|
||||||
confirmButtonText: '确定',
|
confirmButtonText: '确定',
|
||||||
cancelButtonText: '取消',
|
cancelButtonText: '取消',
|
||||||
type: 'warning'
|
type: 'warning'
|
||||||
}).then(() => {
|
}).then(() => {
|
||||||
const ids = this.selectedDictData.map(item => item.id)
|
const ids = selectedRows.map(item => item.id)
|
||||||
dictApi.deleteDictData(ids, ({ data }) => {
|
dictApi.deleteDictData(ids, ({ data }) => {
|
||||||
if (data.code === 0) {
|
if (data.code === 0) {
|
||||||
this.$message.success('删除成功')
|
this.$message.success('删除成功')
|
||||||
@@ -832,4 +834,18 @@ export default {
|
|||||||
flex: 1;
|
flex: 1;
|
||||||
overflow: hidden;
|
overflow: hidden;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
:deep(.el-checkbox__inner) {
|
||||||
|
background-color: #eeeeee !important;
|
||||||
|
border-color: #cccccc !important;
|
||||||
|
}
|
||||||
|
|
||||||
|
:deep(.el-checkbox__inner:hover) {
|
||||||
|
border-color: #cccccc !important;
|
||||||
|
}
|
||||||
|
|
||||||
|
:deep(.el-checkbox__input.is-checked .el-checkbox__inner) {
|
||||||
|
background-color: #5f70f3 !important;
|
||||||
|
border-color: #5f70f3 !important;
|
||||||
|
}
|
||||||
</style>
|
</style>
|
||||||
@@ -318,7 +318,6 @@ export default {
|
|||||||
<style scoped>
|
<style scoped>
|
||||||
.welcome {
|
.welcome {
|
||||||
min-width: 900px;
|
min-width: 900px;
|
||||||
min-height: 506px;
|
|
||||||
height: 100vh;
|
height: 100vh;
|
||||||
display: flex;
|
display: flex;
|
||||||
position: relative;
|
position: relative;
|
||||||
@@ -334,7 +333,7 @@ export default {
|
|||||||
display: flex;
|
display: flex;
|
||||||
justify-content: space-between;
|
justify-content: space-between;
|
||||||
align-items: center;
|
align-items: center;
|
||||||
padding: 16px 24px;
|
padding: 1.5vh 24px;
|
||||||
}
|
}
|
||||||
|
|
||||||
.page-title {
|
.page-title {
|
||||||
@@ -344,7 +343,7 @@ export default {
|
|||||||
}
|
}
|
||||||
|
|
||||||
.main-wrapper {
|
.main-wrapper {
|
||||||
margin: 5px 22px;
|
margin: 1vh 22px;
|
||||||
border-radius: 15px;
|
border-radius: 15px;
|
||||||
height: calc(100vh - 24vh);
|
height: calc(100vh - 24vh);
|
||||||
box-shadow: 0 2px 12px rgba(0, 0, 0, 0.1);
|
box-shadow: 0 2px 12px rgba(0, 0, 0, 0.1);
|
||||||
@@ -416,7 +415,7 @@ export default {
|
|||||||
}
|
}
|
||||||
|
|
||||||
.form-content {
|
.form-content {
|
||||||
padding: 20px 0;
|
padding: 2vh 0;
|
||||||
}
|
}
|
||||||
|
|
||||||
.form-grid {
|
.form-grid {
|
||||||
@@ -450,11 +449,11 @@ export default {
|
|||||||
}
|
}
|
||||||
|
|
||||||
.template-item {
|
.template-item {
|
||||||
height: 37px;
|
height: 4vh;
|
||||||
width: 76px;
|
width: 76px;
|
||||||
border-radius: 8px;
|
border-radius: 8px;
|
||||||
background: #e6ebff;
|
background: #e6ebff;
|
||||||
line-height: 37px;
|
line-height: 4vh;
|
||||||
font-weight: 400;
|
font-weight: 400;
|
||||||
font-size: 11px;
|
font-size: 11px;
|
||||||
text-align: center;
|
text-align: center;
|
||||||
@@ -471,7 +470,7 @@ export default {
|
|||||||
display: flex;
|
display: flex;
|
||||||
flex-wrap: wrap;
|
flex-wrap: wrap;
|
||||||
gap: 8px;
|
gap: 8px;
|
||||||
margin-top: 20px;
|
margin-top: 2vh;
|
||||||
align-items: center;
|
align-items: center;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ from loguru import logger
|
|||||||
from config.config_loader import load_config
|
from config.config_loader import load_config
|
||||||
from config.settings import check_config_file
|
from config.settings import check_config_file
|
||||||
|
|
||||||
SERVER_VERSION = "0.3.14"
|
SERVER_VERSION = "0.4.2"
|
||||||
|
|
||||||
|
|
||||||
def get_module_abbreviation(module_name, module_dict):
|
def get_module_abbreviation(module_name, module_dict):
|
||||||
|
|||||||
@@ -20,6 +20,8 @@ from core.utils.util import (
|
|||||||
get_string_no_punctuation_or_emoji,
|
get_string_no_punctuation_or_emoji,
|
||||||
extract_json_from_string,
|
extract_json_from_string,
|
||||||
initialize_modules,
|
initialize_modules,
|
||||||
|
check_vad_update,
|
||||||
|
check_asr_update,
|
||||||
)
|
)
|
||||||
from concurrent.futures import ThreadPoolExecutor, TimeoutError
|
from concurrent.futures import ThreadPoolExecutor, TimeoutError
|
||||||
from core.handle.sendAudioHandle import sendAudioMessage
|
from core.handle.sendAudioHandle import sendAudioMessage
|
||||||
@@ -54,7 +56,9 @@ class ConnectionHandler:
|
|||||||
_intent,
|
_intent,
|
||||||
server=None,
|
server=None,
|
||||||
):
|
):
|
||||||
self.config = config
|
self.common_config = config
|
||||||
|
self.config = copy.deepcopy(config)
|
||||||
|
self.session_id = str(uuid.uuid4())
|
||||||
self.logger = setup_logging()
|
self.logger = setup_logging()
|
||||||
self.auth = AuthMiddleware(config)
|
self.auth = AuthMiddleware(config)
|
||||||
self.server = server # 保存server实例的引用
|
self.server = server # 保存server实例的引用
|
||||||
@@ -68,7 +72,6 @@ class ConnectionHandler:
|
|||||||
self.device_id = None
|
self.device_id = None
|
||||||
self.client_ip = None
|
self.client_ip = None
|
||||||
self.client_ip_info = {}
|
self.client_ip_info = {}
|
||||||
self.session_id = None
|
|
||||||
self.prompt = None
|
self.prompt = None
|
||||||
self.welcome_msg = None
|
self.welcome_msg = None
|
||||||
self.max_output_size = 0
|
self.max_output_size = 0
|
||||||
@@ -89,8 +92,10 @@ class ConnectionHandler:
|
|||||||
self.tts_report_thread = None
|
self.tts_report_thread = None
|
||||||
|
|
||||||
# 依赖的组件
|
# 依赖的组件
|
||||||
self.vad = _vad
|
self.vad = None
|
||||||
self.asr = _asr
|
self.asr = None
|
||||||
|
self._asr = _asr
|
||||||
|
self._vad = _vad
|
||||||
self.llm = _llm
|
self.llm = _llm
|
||||||
self.tts = _tts
|
self.tts = _tts
|
||||||
self.memory = _memory
|
self.memory = _memory
|
||||||
@@ -133,6 +138,8 @@ class ConnectionHandler:
|
|||||||
int(self.config.get("close_connection_no_voice_time", 120)) + 60
|
int(self.config.get("close_connection_no_voice_time", 120)) + 60
|
||||||
) # 在原来第一道关闭的基础上加60秒,进行二道关闭
|
) # 在原来第一道关闭的基础上加60秒,进行二道关闭
|
||||||
|
|
||||||
|
self.audio_format = "opus"
|
||||||
|
|
||||||
async def handle_connection(self, ws):
|
async def handle_connection(self, ws):
|
||||||
try:
|
try:
|
||||||
# 获取并验证headers
|
# 获取并验证headers
|
||||||
@@ -168,7 +175,6 @@ class ConnectionHandler:
|
|||||||
# 认证通过,继续处理
|
# 认证通过,继续处理
|
||||||
self.websocket = ws
|
self.websocket = ws
|
||||||
self.device_id = self.headers.get("device-id", None)
|
self.device_id = self.headers.get("device-id", None)
|
||||||
self.session_id = str(uuid.uuid4())
|
|
||||||
|
|
||||||
# 启动超时检查任务
|
# 启动超时检查任务
|
||||||
self.timeout_task = asyncio.create_task(self._check_timeout())
|
self.timeout_task = asyncio.create_task(self._check_timeout())
|
||||||
@@ -178,9 +184,9 @@ class ConnectionHandler:
|
|||||||
await self.websocket.send(json.dumps(self.welcome_msg))
|
await self.websocket.send(json.dumps(self.welcome_msg))
|
||||||
|
|
||||||
# 获取差异化配置
|
# 获取差异化配置
|
||||||
private_config = self._initialize_private_config()
|
self._initialize_private_config()
|
||||||
# 异步初始化
|
# 异步初始化
|
||||||
self.executor.submit(self._initialize_components, private_config)
|
self.executor.submit(self._initialize_components)
|
||||||
# tts 消化线程
|
# tts 消化线程
|
||||||
self.tts_priority_thread = threading.Thread(
|
self.tts_priority_thread = threading.Thread(
|
||||||
target=self._tts_priority_thread, daemon=True
|
target=self._tts_priority_thread, daemon=True
|
||||||
@@ -273,13 +279,20 @@ class ConnectionHandler:
|
|||||||
"message": f"Restart failed: {str(e)}"
|
"message": f"Restart failed: {str(e)}"
|
||||||
}))
|
}))
|
||||||
|
|
||||||
def _initialize_components(self, private_config):
|
def _initialize_components(self):
|
||||||
"""初始化组件"""
|
"""初始化组件"""
|
||||||
if private_config is not None:
|
if self.config.get("prompt") is not None:
|
||||||
self._initialize_models(private_config)
|
|
||||||
else:
|
|
||||||
self.prompt = self.config["prompt"]
|
self.prompt = self.config["prompt"]
|
||||||
self.change_system_prompt(self.prompt)
|
self.change_system_prompt(self.prompt)
|
||||||
|
self.logger.bind(tag=TAG).info(
|
||||||
|
f"初始化组件: prompt成功 {self.prompt[:50]}..."
|
||||||
|
)
|
||||||
|
|
||||||
|
"""初始化本地组件"""
|
||||||
|
if self.vad is None:
|
||||||
|
self.vad = self._vad
|
||||||
|
if self.asr is None:
|
||||||
|
self.asr = self._asr
|
||||||
"""加载记忆"""
|
"""加载记忆"""
|
||||||
self._initialize_memory()
|
self._initialize_memory()
|
||||||
"""加载意图识别"""
|
"""加载意图识别"""
|
||||||
@@ -326,54 +339,21 @@ class ConnectionHandler:
|
|||||||
self.logger.bind(tag=TAG).error(f"获取差异化配置失败: {e}")
|
self.logger.bind(tag=TAG).error(f"获取差异化配置失败: {e}")
|
||||||
private_config = {}
|
private_config = {}
|
||||||
|
|
||||||
init_tts = False
|
init_llm, init_tts, init_memory, init_intent = (
|
||||||
if private_config.get("TTS", None) is not None:
|
|
||||||
init_tts = True
|
|
||||||
self.config["TTS"] = private_config["TTS"]
|
|
||||||
self.config["selected_module"]["TTS"] = private_config["selected_module"][
|
|
||||||
"TTS"
|
|
||||||
]
|
|
||||||
|
|
||||||
try:
|
|
||||||
modules = initialize_modules(
|
|
||||||
self.logger,
|
|
||||||
private_config,
|
|
||||||
False,
|
|
||||||
False,
|
|
||||||
False,
|
|
||||||
init_tts,
|
|
||||||
False,
|
|
||||||
False,
|
|
||||||
)
|
|
||||||
except Exception as e:
|
|
||||||
self.logger.bind(tag=TAG).error(f"初始化组件失败: {e}")
|
|
||||||
modules = {}
|
|
||||||
if modules.get("tts", None) is not None:
|
|
||||||
self.tts = modules["tts"]
|
|
||||||
if modules.get("prompt", None) is not None:
|
|
||||||
self.change_system_prompt(modules["prompt"])
|
|
||||||
private_config["prompt"] = None
|
|
||||||
return private_config
|
|
||||||
|
|
||||||
def _initialize_models(self, private_config):
|
|
||||||
init_vad, init_asr, init_llm, init_memory, init_intent = (
|
|
||||||
False,
|
|
||||||
False,
|
False,
|
||||||
False,
|
False,
|
||||||
False,
|
False,
|
||||||
False,
|
False,
|
||||||
)
|
)
|
||||||
if private_config.get("VAD", None) is not None:
|
|
||||||
init_vad = True
|
init_vad = check_vad_update(self.common_config, private_config)
|
||||||
self.config["VAD"] = private_config["VAD"]
|
init_asr = check_asr_update(self.common_config, private_config)
|
||||||
self.config["selected_module"]["VAD"] = private_config["selected_module"][
|
|
||||||
"VAD"
|
if private_config.get("TTS", None) is not None:
|
||||||
]
|
init_tts = True
|
||||||
if private_config.get("ASR", None) is not None:
|
self.config["TTS"] = private_config["TTS"]
|
||||||
init_asr = True
|
self.config["selected_module"]["TTS"] = private_config["selected_module"][
|
||||||
self.config["ASR"] = private_config["ASR"]
|
"TTS"
|
||||||
self.config["selected_module"]["ASR"] = private_config["selected_module"][
|
|
||||||
"ASR"
|
|
||||||
]
|
]
|
||||||
if private_config.get("LLM", None) is not None:
|
if private_config.get("LLM", None) is not None:
|
||||||
init_llm = True
|
init_llm = True
|
||||||
@@ -393,8 +373,11 @@ class ConnectionHandler:
|
|||||||
self.config["selected_module"]["Intent"] = private_config[
|
self.config["selected_module"]["Intent"] = private_config[
|
||||||
"selected_module"
|
"selected_module"
|
||||||
]["Intent"]
|
]["Intent"]
|
||||||
|
if private_config.get("prompt", None) is not None:
|
||||||
|
self.config["prompt"] = private_config["prompt"]
|
||||||
if private_config.get("device_max_output_size", None) is not None:
|
if private_config.get("device_max_output_size", None) is not None:
|
||||||
self.max_output_size = int(private_config["device_max_output_size"])
|
self.max_output_size = int(private_config["device_max_output_size"])
|
||||||
|
|
||||||
try:
|
try:
|
||||||
modules = initialize_modules(
|
modules = initialize_modules(
|
||||||
self.logger,
|
self.logger,
|
||||||
@@ -402,13 +385,15 @@ class ConnectionHandler:
|
|||||||
init_vad,
|
init_vad,
|
||||||
init_asr,
|
init_asr,
|
||||||
init_llm,
|
init_llm,
|
||||||
False,
|
init_tts,
|
||||||
init_memory,
|
init_memory,
|
||||||
init_intent,
|
init_intent,
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
self.logger.bind(tag=TAG).error(f"初始化组件失败: {e}")
|
self.logger.bind(tag=TAG).error(f"初始化组件失败: {e}")
|
||||||
modules = {}
|
modules = {}
|
||||||
|
if modules.get("tts", None) is not None:
|
||||||
|
self.tts = modules["tts"]
|
||||||
if modules.get("vad", None) is not None:
|
if modules.get("vad", None) is not None:
|
||||||
self.vad = modules["vad"]
|
self.vad = modules["vad"]
|
||||||
if modules.get("asr", None) is not None:
|
if modules.get("asr", None) is not None:
|
||||||
@@ -486,10 +471,12 @@ class ConnectionHandler:
|
|||||||
processed_chars = 0 # 跟踪已处理的字符位置
|
processed_chars = 0 # 跟踪已处理的字符位置
|
||||||
try:
|
try:
|
||||||
# 使用带记忆的对话
|
# 使用带记忆的对话
|
||||||
future = asyncio.run_coroutine_threadsafe(
|
memory_str = None
|
||||||
self.memory.query_memory(query), self.loop
|
if self.memory is not None:
|
||||||
)
|
future = asyncio.run_coroutine_threadsafe(
|
||||||
memory_str = future.result()
|
self.memory.query_memory(query), self.loop
|
||||||
|
)
|
||||||
|
memory_str = future.result()
|
||||||
|
|
||||||
self.logger.bind(tag=TAG).debug(f"记忆内容: {memory_str}")
|
self.logger.bind(tag=TAG).debug(f"记忆内容: {memory_str}")
|
||||||
llm_responses = self.llm.response(
|
llm_responses = self.llm.response(
|
||||||
@@ -550,7 +537,7 @@ class ConnectionHandler:
|
|||||||
future = self.executor.submit(
|
future = self.executor.submit(
|
||||||
self.speak_and_play, segment_text, text_index
|
self.speak_and_play, segment_text, text_index
|
||||||
)
|
)
|
||||||
self.tts_queue.put(future)
|
self.tts_queue.put((future, text_index))
|
||||||
|
|
||||||
self.llm_finish_task = True
|
self.llm_finish_task = True
|
||||||
self.dialogue.put(Message(role="assistant", content="".join(response_message)))
|
self.dialogue.put(Message(role="assistant", content="".join(response_message)))
|
||||||
@@ -577,10 +564,12 @@ class ConnectionHandler:
|
|||||||
start_time = time.time()
|
start_time = time.time()
|
||||||
|
|
||||||
# 使用带记忆的对话
|
# 使用带记忆的对话
|
||||||
future = asyncio.run_coroutine_threadsafe(
|
memory_str = None
|
||||||
self.memory.query_memory(query), self.loop
|
if self.memory is not None:
|
||||||
)
|
future = asyncio.run_coroutine_threadsafe(
|
||||||
memory_str = future.result()
|
self.memory.query_memory(query), self.loop
|
||||||
|
)
|
||||||
|
memory_str = future.result()
|
||||||
|
|
||||||
# self.logger.bind(tag=TAG).info(f"对话记录: {self.dialogue.get_llm_dialogue_with_memory(memory_str)}")
|
# self.logger.bind(tag=TAG).info(f"对话记录: {self.dialogue.get_llm_dialogue_with_memory(memory_str)}")
|
||||||
|
|
||||||
@@ -725,7 +714,7 @@ class ConnectionHandler:
|
|||||||
future = self.executor.submit(
|
future = self.executor.submit(
|
||||||
self.speak_and_play, segment_text, text_index
|
self.speak_and_play, segment_text, text_index
|
||||||
)
|
)
|
||||||
self.tts_queue.put(future)
|
self.tts_queue.put((future, text_index))
|
||||||
|
|
||||||
# 存储对话内容
|
# 存储对话内容
|
||||||
if len(response_message) > 0:
|
if len(response_message) > 0:
|
||||||
@@ -787,7 +776,7 @@ class ConnectionHandler:
|
|||||||
text = result.response
|
text = result.response
|
||||||
self.recode_first_last_text(text, text_index)
|
self.recode_first_last_text(text, text_index)
|
||||||
future = self.executor.submit(self.speak_and_play, text, text_index)
|
future = self.executor.submit(self.speak_and_play, text, text_index)
|
||||||
self.tts_queue.put(future)
|
self.tts_queue.put((future, text_index))
|
||||||
self.dialogue.put(Message(role="assistant", content=text))
|
self.dialogue.put(Message(role="assistant", content=text))
|
||||||
elif result.action == Action.REQLLM: # 调用函数后再请求llm生成回复
|
elif result.action == Action.REQLLM: # 调用函数后再请求llm生成回复
|
||||||
text = result.result
|
text = result.result
|
||||||
@@ -820,7 +809,7 @@ class ConnectionHandler:
|
|||||||
text = result.result
|
text = result.result
|
||||||
self.recode_first_last_text(text, text_index)
|
self.recode_first_last_text(text, text_index)
|
||||||
future = self.executor.submit(self.speak_and_play, text, text_index)
|
future = self.executor.submit(self.speak_and_play, text, text_index)
|
||||||
self.tts_queue.put(future)
|
self.tts_queue.put((future, text_index))
|
||||||
self.dialogue.put(Message(role="assistant", content=text))
|
self.dialogue.put(Message(role="assistant", content=text))
|
||||||
else:
|
else:
|
||||||
pass
|
pass
|
||||||
@@ -841,7 +830,7 @@ class ConnectionHandler:
|
|||||||
if future is None:
|
if future is None:
|
||||||
continue
|
continue
|
||||||
text = None
|
text = None
|
||||||
opus_datas, tts_file = [], None
|
opus_datas, tts_file = [], None
|
||||||
try:
|
try:
|
||||||
self.logger.bind(tag=TAG).debug("正在处理TTS任务...")
|
self.logger.bind(tag=TAG).debug("正在处理TTS任务...")
|
||||||
tts_timeout = int(self.config.get("tts_timeout", 10))
|
tts_timeout = int(self.config.get("tts_timeout", 10))
|
||||||
@@ -859,9 +848,12 @@ class ConnectionHandler:
|
|||||||
f"TTS生成:文件路径: {tts_file}"
|
f"TTS生成:文件路径: {tts_file}"
|
||||||
)
|
)
|
||||||
if os.path.exists(tts_file):
|
if os.path.exists(tts_file):
|
||||||
opus_datas, _ = self.tts.audio_to_opus_data(tts_file)
|
if self.audio_format == "pcm":
|
||||||
|
audio_datas, _ = self.tts.audio_to_pcm_data(tts_file)
|
||||||
|
else:
|
||||||
|
audio_datas, _ = self.tts.audio_to_opus_data(tts_file)
|
||||||
# 在这里上报TTS数据(使用文件路径)
|
# 在这里上报TTS数据(使用文件路径)
|
||||||
enqueue_tts_report(self, 2, text, opus_datas)
|
enqueue_tts_report(self, 2, text, audio_datas)
|
||||||
else:
|
else:
|
||||||
self.logger.bind(tag=TAG).error(
|
self.logger.bind(tag=TAG).error(
|
||||||
f"TTS出错:文件不存在{tts_file}"
|
f"TTS出错:文件不存在{tts_file}"
|
||||||
@@ -872,7 +864,7 @@ class ConnectionHandler:
|
|||||||
self.logger.bind(tag=TAG).error(f"TTS出错: {e}")
|
self.logger.bind(tag=TAG).error(f"TTS出错: {e}")
|
||||||
if not self.client_abort:
|
if not self.client_abort:
|
||||||
# 如果没有中途打断就发送语音
|
# 如果没有中途打断就发送语音
|
||||||
self.audio_play_queue.put((opus_datas, text, text_index))
|
self.audio_play_queue.put((audio_datas, text, text_index))
|
||||||
if (
|
if (
|
||||||
self.tts.delete_audio_file
|
self.tts.delete_audio_file
|
||||||
and tts_file is not None
|
and tts_file is not None
|
||||||
@@ -903,13 +895,13 @@ class ConnectionHandler:
|
|||||||
text = None
|
text = None
|
||||||
try:
|
try:
|
||||||
try:
|
try:
|
||||||
opus_datas, text, text_index = self.audio_play_queue.get(timeout=1)
|
audio_datas, text, text_index = self.audio_play_queue.get(timeout=1)
|
||||||
except queue.Empty:
|
except queue.Empty:
|
||||||
if self.stop_event.is_set():
|
if self.stop_event.is_set():
|
||||||
break
|
break
|
||||||
continue
|
continue
|
||||||
future = asyncio.run_coroutine_threadsafe(
|
future = asyncio.run_coroutine_threadsafe(
|
||||||
sendAudioMessage(self, opus_datas, text, text_index), self.loop
|
sendAudioMessage(self, audio_datas, text, text_index), self.loop
|
||||||
)
|
)
|
||||||
future.result()
|
future.result()
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
|||||||
@@ -3,15 +3,16 @@ import queue
|
|||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
|
||||||
|
|
||||||
|
|
||||||
async def handleAbortMessage(conn):
|
async def handleAbortMessage(conn):
|
||||||
logger.bind(tag=TAG).info("Abort message received")
|
conn.logger.bind(tag=TAG).info("Abort message received")
|
||||||
# 设置成打断状态,会自动打断llm、tts任务
|
# 设置成打断状态,会自动打断llm、tts任务
|
||||||
conn.client_abort = True
|
conn.client_abort = True
|
||||||
conn.clear_queues()
|
conn.clear_queues()
|
||||||
# 打断客户端说话状态
|
# 打断客户端说话状态
|
||||||
await conn.websocket.send(json.dumps({"type": "tts", "state": "stop", "session_id": conn.session_id}))
|
await conn.websocket.send(
|
||||||
|
json.dumps({"type": "tts", "state": "stop", "session_id": conn.session_id})
|
||||||
|
)
|
||||||
conn.clearSpeakStatus()
|
conn.clearSpeakStatus()
|
||||||
logger.bind(tag=TAG).info("Abort message received-end")
|
conn.logger.bind(tag=TAG).info("Abort message received-end")
|
||||||
|
|||||||
@@ -4,7 +4,6 @@ from plugins_func.register import FunctionRegistry, ActionResponse, Action, Tool
|
|||||||
from plugins_func.functions.hass_init import append_devices_to_prompt
|
from plugins_func.functions.hass_init import append_devices_to_prompt
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
|
||||||
|
|
||||||
|
|
||||||
class FunctionHandler:
|
class FunctionHandler:
|
||||||
@@ -40,7 +39,9 @@ class FunctionHandler:
|
|||||||
for func in self.functions_desc:
|
for func in self.functions_desc:
|
||||||
func_names.append(func["function"]["name"])
|
func_names.append(func["function"]["name"])
|
||||||
# 打印当前支持的函数列表
|
# 打印当前支持的函数列表
|
||||||
logger.bind(tag=TAG).info(f"当前支持的函数列表: {func_names}")
|
self.conn.logger.bind(tag=TAG, session_id=self.conn.session_id).info(
|
||||||
|
f"当前支持的函数列表: {func_names}"
|
||||||
|
)
|
||||||
return func_names
|
return func_names
|
||||||
|
|
||||||
def get_functions(self):
|
def get_functions(self):
|
||||||
@@ -79,7 +80,9 @@ class FunctionHandler:
|
|||||||
func = funcItem.func
|
func = funcItem.func
|
||||||
arguments = function_call_data["arguments"]
|
arguments = function_call_data["arguments"]
|
||||||
arguments = json.loads(arguments) if arguments else {}
|
arguments = json.loads(arguments) if arguments else {}
|
||||||
logger.bind(tag=TAG).debug(f"调用函数: {function_name}, 参数: {arguments}")
|
self.conn.logger.bind(tag=TAG).debug(
|
||||||
|
f"调用函数: {function_name}, 参数: {arguments}"
|
||||||
|
)
|
||||||
if (
|
if (
|
||||||
funcItem.type == ToolType.SYSTEM_CTL
|
funcItem.type == ToolType.SYSTEM_CTL
|
||||||
or funcItem.type == ToolType.IOT_CTL
|
or funcItem.type == ToolType.IOT_CTL
|
||||||
@@ -94,6 +97,6 @@ class FunctionHandler:
|
|||||||
action=Action.NOTFOUND, result="没有找到对应的函数", response=""
|
action=Action.NOTFOUND, result="没有找到对应的函数", response=""
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"处理function call错误: {e}")
|
self.conn.logger.bind(tag=TAG).error(f"处理function call错误: {e}")
|
||||||
|
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -1,5 +1,4 @@
|
|||||||
import json
|
import json
|
||||||
from config.logger import setup_logging
|
|
||||||
from core.handle.sendAudioHandle import send_stt_message
|
from core.handle.sendAudioHandle import send_stt_message
|
||||||
from core.utils.util import remove_punctuation_and_length
|
from core.utils.util import remove_punctuation_and_length
|
||||||
import shutil
|
import shutil
|
||||||
@@ -9,7 +8,6 @@ import random
|
|||||||
import time
|
import time
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
|
||||||
|
|
||||||
WAKEUP_CONFIG = {
|
WAKEUP_CONFIG = {
|
||||||
"dir": "config/assets/",
|
"dir": "config/assets/",
|
||||||
@@ -21,7 +19,16 @@ WAKEUP_CONFIG = {
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
async def handleHelloMessage(conn):
|
async def handleHelloMessage(conn, msg_json):
|
||||||
|
"""处理hello消息"""
|
||||||
|
audio_params = msg_json.get("audio_params")
|
||||||
|
if audio_params:
|
||||||
|
format = audio_params.get("format")
|
||||||
|
conn.logger.bind(tag=TAG).info(f"客户端音频格式: {format}")
|
||||||
|
conn.audio_format = format
|
||||||
|
conn.asr.set_audio_format(format)
|
||||||
|
conn.welcome_msg["audio_params"] = audio_params
|
||||||
|
|
||||||
await conn.websocket.send(json.dumps(conn.welcome_msg))
|
await conn.websocket.send(json.dumps(conn.welcome_msg))
|
||||||
|
|
||||||
|
|
||||||
@@ -75,7 +82,7 @@ async def wakeupWordsResponse(conn):
|
|||||||
await asyncio.sleep(1)
|
await asyncio.sleep(1)
|
||||||
wait_max_time -= 1
|
wait_max_time -= 1
|
||||||
if wait_max_time <= 0:
|
if wait_max_time <= 0:
|
||||||
logger.bind(tag=TAG).error("连接对象没有llm")
|
conn.logger.bind(tag=TAG).error("连接对象没有llm")
|
||||||
return
|
return
|
||||||
|
|
||||||
"""唤醒词响应"""
|
"""唤醒词响应"""
|
||||||
|
|||||||
@@ -5,10 +5,8 @@ from core.handle.sendAudioHandle import send_stt_message
|
|||||||
from core.handle.helloHandle import checkWakeupWords
|
from core.handle.helloHandle import checkWakeupWords
|
||||||
from core.utils.util import remove_punctuation_and_length
|
from core.utils.util import remove_punctuation_and_length
|
||||||
from core.utils.dialogue import Message
|
from core.utils.dialogue import Message
|
||||||
from loguru import logger
|
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
|
||||||
|
|
||||||
|
|
||||||
async def handle_user_intent(conn, text):
|
async def handle_user_intent(conn, text):
|
||||||
@@ -36,7 +34,7 @@ async def check_direct_exit(conn, text):
|
|||||||
cmd_exit = conn.cmd_exit
|
cmd_exit = conn.cmd_exit
|
||||||
for cmd in cmd_exit:
|
for cmd in cmd_exit:
|
||||||
if text == cmd:
|
if text == cmd:
|
||||||
logger.bind(tag=TAG).info(f"识别到明确的退出命令: {text}")
|
conn.logger.bind(tag=TAG).info(f"识别到明确的退出命令: {text}")
|
||||||
await send_stt_message(conn, text)
|
await send_stt_message(conn, text)
|
||||||
await conn.close()
|
await conn.close()
|
||||||
return True
|
return True
|
||||||
@@ -46,7 +44,7 @@ async def check_direct_exit(conn, text):
|
|||||||
async def analyze_intent_with_llm(conn, text):
|
async def analyze_intent_with_llm(conn, text):
|
||||||
"""使用LLM分析用户意图"""
|
"""使用LLM分析用户意图"""
|
||||||
if not hasattr(conn, "intent") or not conn.intent:
|
if not hasattr(conn, "intent") or not conn.intent:
|
||||||
logger.bind(tag=TAG).warning("意图识别服务未初始化")
|
conn.logger.bind(tag=TAG).warning("意图识别服务未初始化")
|
||||||
return None
|
return None
|
||||||
|
|
||||||
# 对话历史记录
|
# 对话历史记录
|
||||||
@@ -55,7 +53,7 @@ async def analyze_intent_with_llm(conn, text):
|
|||||||
intent_result = await conn.intent.detect_intent(conn, dialogue.dialogue, text)
|
intent_result = await conn.intent.detect_intent(conn, dialogue.dialogue, text)
|
||||||
return intent_result
|
return intent_result
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"意图识别失败: {str(e)}")
|
conn.logger.bind(tag=TAG).error(f"意图识别失败: {str(e)}")
|
||||||
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -69,7 +67,7 @@ async def process_intent_result(conn, intent_result, original_text):
|
|||||||
# 检查是否有function_call
|
# 检查是否有function_call
|
||||||
if "function_call" in intent_data:
|
if "function_call" in intent_data:
|
||||||
# 直接从意图识别获取了function_call
|
# 直接从意图识别获取了function_call
|
||||||
logger.bind(tag=TAG).debug(
|
conn.logger.bind(tag=TAG).debug(
|
||||||
f"检测到function_call格式的意图结果: {intent_data['function_call']['name']}"
|
f"检测到function_call格式的意图结果: {intent_data['function_call']['name']}"
|
||||||
)
|
)
|
||||||
function_name = intent_data["function_call"]["name"]
|
function_name = intent_data["function_call"]["name"]
|
||||||
@@ -118,7 +116,7 @@ async def process_intent_result(conn, intent_result, original_text):
|
|||||||
conn.speak_and_play, text, text_index
|
conn.speak_and_play, text, text_index
|
||||||
)
|
)
|
||||||
conn.llm_finish_task = True
|
conn.llm_finish_task = True
|
||||||
conn.tts_queue.put(future)
|
conn.tts_queue.put((future, text_index))
|
||||||
conn.dialogue.put(Message(role="assistant", content=text))
|
conn.dialogue.put(Message(role="assistant", content=text))
|
||||||
|
|
||||||
# 将函数执行放在线程池中
|
# 将函数执行放在线程池中
|
||||||
@@ -126,7 +124,7 @@ async def process_intent_result(conn, intent_result, original_text):
|
|||||||
return True
|
return True
|
||||||
return False
|
return False
|
||||||
except json.JSONDecodeError as e:
|
except json.JSONDecodeError as e:
|
||||||
logger.bind(tag=TAG).error(f"处理意图结果时出错: {e}")
|
conn.logger.bind(tag=TAG).error(f"处理意图结果时出错: {e}")
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -10,7 +10,6 @@ from plugins_func.register import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
|
||||||
|
|
||||||
|
|
||||||
def wrap_async_function(async_func):
|
def wrap_async_function(async_func):
|
||||||
@@ -21,7 +20,7 @@ def wrap_async_function(async_func):
|
|||||||
# 获取连接对象(第一个参数)
|
# 获取连接对象(第一个参数)
|
||||||
conn = args[0]
|
conn = args[0]
|
||||||
if not hasattr(conn, "loop"):
|
if not hasattr(conn, "loop"):
|
||||||
logger.bind(tag=TAG).error("Connection对象没有loop属性")
|
conn.logger.bind(tag=TAG).error("Connection对象没有loop属性")
|
||||||
return ActionResponse(
|
return ActionResponse(
|
||||||
Action.ERROR,
|
Action.ERROR,
|
||||||
"Connection对象没有loop属性",
|
"Connection对象没有loop属性",
|
||||||
@@ -35,7 +34,7 @@ def wrap_async_function(async_func):
|
|||||||
# 等待结果返回
|
# 等待结果返回
|
||||||
return future.result()
|
return future.result()
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"运行异步函数时出错: {e}")
|
conn.logger.bind(tag=TAG).error(f"运行异步函数时出错: {e}")
|
||||||
return ActionResponse(Action.ERROR, str(e), f"执行操作时出错: {e}")
|
return ActionResponse(Action.ERROR, str(e), f"执行操作时出错: {e}")
|
||||||
|
|
||||||
return wrapper
|
return wrapper
|
||||||
@@ -57,7 +56,7 @@ def create_iot_function(device_name, method_name, method_info):
|
|||||||
response_failure = "操作失败"
|
response_failure = "操作失败"
|
||||||
|
|
||||||
# 打印响应参数
|
# 打印响应参数
|
||||||
logger.bind(tag=TAG).debug(
|
conn.logger.bind(tag=TAG).debug(
|
||||||
f"控制函数接收到的响应参数: success='{response_success}', failure='{response_failure}'"
|
f"控制函数接收到的响应参数: success='{response_success}', failure='{response_failure}'"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -86,7 +85,9 @@ def create_iot_function(device_name, method_name, method_info):
|
|||||||
|
|
||||||
return ActionResponse(Action.RESPONSE, result, response)
|
return ActionResponse(Action.RESPONSE, result, response)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"执行{device_name}的{method_name}操作失败: {e}")
|
conn.logger.bind(tag=TAG).error(
|
||||||
|
f"执行{device_name}的{method_name}操作失败: {e}"
|
||||||
|
)
|
||||||
|
|
||||||
# 操作失败时使用大模型提供的失败响应
|
# 操作失败时使用大模型提供的失败响应
|
||||||
response = response_failure
|
response = response_failure
|
||||||
@@ -104,7 +105,7 @@ def create_iot_query_function(device_name, prop_name, prop_info):
|
|||||||
async def iot_query_function(conn, response_success=None, response_failure=None):
|
async def iot_query_function(conn, response_success=None, response_failure=None):
|
||||||
try:
|
try:
|
||||||
# 打印响应参数
|
# 打印响应参数
|
||||||
logger.bind(tag=TAG).info(
|
conn.logger.bind(tag=TAG).info(
|
||||||
f"查询函数接收到的响应参数: success='{response_success}', failure='{response_failure}'"
|
f"查询函数接收到的响应参数: success='{response_success}', failure='{response_failure}'"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -122,7 +123,9 @@ def create_iot_query_function(device_name, prop_name, prop_info):
|
|||||||
|
|
||||||
return ActionResponse(Action.ERROR, f"属性{prop_name}不存在", response)
|
return ActionResponse(Action.ERROR, f"属性{prop_name}不存在", response)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"查询{device_name}的{prop_name}时出错: {e}")
|
conn.logger.bind(tag=TAG).error(
|
||||||
|
f"查询{device_name}的{prop_name}时出错: {e}"
|
||||||
|
)
|
||||||
|
|
||||||
# 查询出错时使用大模型提供的失败响应
|
# 查询出错时使用大模型提供的失败响应
|
||||||
response = response_failure
|
response = response_failure
|
||||||
@@ -280,7 +283,7 @@ async def handleIotDescriptors(conn, descriptors):
|
|||||||
await asyncio.sleep(1)
|
await asyncio.sleep(1)
|
||||||
wait_max_time -= 1
|
wait_max_time -= 1
|
||||||
if wait_max_time <= 0:
|
if wait_max_time <= 0:
|
||||||
logger.bind(tag=TAG).debug("连接对象没有func_handler")
|
conn.logger.bind(tag=TAG).debug("连接对象没有func_handler")
|
||||||
return
|
return
|
||||||
"""处理物联网描述"""
|
"""处理物联网描述"""
|
||||||
functions_changed = False
|
functions_changed = False
|
||||||
@@ -323,7 +326,7 @@ async def handleIotDescriptors(conn, descriptors):
|
|||||||
if hasattr(conn, "func_handler"):
|
if hasattr(conn, "func_handler"):
|
||||||
for func_name in device_functions:
|
for func_name in device_functions:
|
||||||
conn.func_handler.function_registry.register_function(func_name)
|
conn.func_handler.function_registry.register_function(func_name)
|
||||||
logger.bind(tag=TAG).info(
|
conn.logger.bind(tag=TAG).info(
|
||||||
f"注册IOT函数到function handler: {func_name}"
|
f"注册IOT函数到function handler: {func_name}"
|
||||||
)
|
)
|
||||||
functions_changed = True
|
functions_changed = True
|
||||||
@@ -332,8 +335,8 @@ async def handleIotDescriptors(conn, descriptors):
|
|||||||
if functions_changed and hasattr(conn, "func_handler"):
|
if functions_changed and hasattr(conn, "func_handler"):
|
||||||
conn.func_handler.upload_functions_desc()
|
conn.func_handler.upload_functions_desc()
|
||||||
func_names = conn.func_handler.current_support_functions()
|
func_names = conn.func_handler.current_support_functions()
|
||||||
logger.bind(tag=TAG).info(f"设备类型: {type_id}")
|
conn.logger.bind(tag=TAG).info(f"设备类型: {type_id}")
|
||||||
logger.bind(tag=TAG).info(
|
conn.logger.bind(tag=TAG).info(
|
||||||
f"更新function描述列表完成,当前支持的函数: {func_names}"
|
f"更新function描述列表完成,当前支持的函数: {func_names}"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -347,13 +350,13 @@ async def handleIotStatus(conn, states):
|
|||||||
for k, v in state["state"].items():
|
for k, v in state["state"].items():
|
||||||
if property_item["name"] == k:
|
if property_item["name"] == k:
|
||||||
if type(v) != type(property_item["value"]):
|
if type(v) != type(property_item["value"]):
|
||||||
logger.bind(tag=TAG).error(
|
conn.logger.bind(tag=TAG).error(
|
||||||
f"属性{property_item['name']}的值类型不匹配"
|
f"属性{property_item['name']}的值类型不匹配"
|
||||||
)
|
)
|
||||||
break
|
break
|
||||||
else:
|
else:
|
||||||
property_item["value"] = v
|
property_item["value"] = v
|
||||||
logger.bind(tag=TAG).info(
|
conn.logger.bind(tag=TAG).info(
|
||||||
f"物联网状态更新: {key} , {property_item['name']} = {v}"
|
f"物联网状态更新: {key} , {property_item['name']} = {v}"
|
||||||
)
|
)
|
||||||
break
|
break
|
||||||
@@ -367,7 +370,7 @@ async def get_iot_status(conn, name, property_name):
|
|||||||
for property_item in value.properties:
|
for property_item in value.properties:
|
||||||
if property_item["name"] == property_name:
|
if property_item["name"] == property_name:
|
||||||
return property_item["value"]
|
return property_item["value"]
|
||||||
logger.bind(tag=TAG).warning(f"未找到设备 {name} 的属性 {property_name}")
|
conn.logger.bind(tag=TAG).warning(f"未找到设备 {name} 的属性 {property_name}")
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
@@ -378,16 +381,16 @@ async def set_iot_status(conn, name, property_name, value):
|
|||||||
for property_item in iot_descriptor.properties:
|
for property_item in iot_descriptor.properties:
|
||||||
if property_item["name"] == property_name:
|
if property_item["name"] == property_name:
|
||||||
if type(value) != type(property_item["value"]):
|
if type(value) != type(property_item["value"]):
|
||||||
logger.bind(tag=TAG).error(
|
conn.logger.bind(tag=TAG).error(
|
||||||
f"属性{property_item['name']}的值类型不匹配"
|
f"属性{property_item['name']}的值类型不匹配"
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
property_item["value"] = value
|
property_item["value"] = value
|
||||||
logger.bind(tag=TAG).info(
|
conn.logger.bind(tag=TAG).info(
|
||||||
f"物联网状态更新: {name} , {property_name} = {value}"
|
f"物联网状态更新: {name} , {property_name} = {value}"
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
logger.bind(tag=TAG).warning(f"未找到设备 {name} 的属性 {property_name}")
|
conn.logger.bind(tag=TAG).warning(f"未找到设备 {name} 的属性 {property_name}")
|
||||||
|
|
||||||
|
|
||||||
async def send_iot_conn(conn, name, method_name, parameters):
|
async def send_iot_conn(conn, name, method_name, parameters):
|
||||||
@@ -409,6 +412,6 @@ async def send_iot_conn(conn, name, method_name, parameters):
|
|||||||
command["parameters"] = parameters
|
command["parameters"] = parameters
|
||||||
send_message = json.dumps({"type": "iot", "commands": [command]})
|
send_message = json.dumps({"type": "iot", "commands": [command]})
|
||||||
await conn.websocket.send(send_message)
|
await conn.websocket.send(send_message)
|
||||||
logger.bind(tag=TAG).info(f"发送物联网指令: {send_message}")
|
conn.logger.bind(tag=TAG).info(f"发送物联网指令: {send_message}")
|
||||||
return
|
return
|
||||||
logger.bind(tag=TAG).error(f"未找到方法{method_name}")
|
conn.logger.bind(tag=TAG).error(f"未找到方法{method_name}")
|
||||||
|
|||||||
@@ -1,4 +1,3 @@
|
|||||||
from config.logger import setup_logging
|
|
||||||
import time
|
import time
|
||||||
import copy
|
import copy
|
||||||
from core.utils.util import remove_punctuation_and_length
|
from core.utils.util import remove_punctuation_and_length
|
||||||
@@ -6,15 +5,16 @@ from core.handle.sendAudioHandle import send_stt_message
|
|||||||
from core.handle.intentHandler import handle_user_intent
|
from core.handle.intentHandler import handle_user_intent
|
||||||
from core.utils.output_counter import check_device_output_limit
|
from core.utils.output_counter import check_device_output_limit
|
||||||
from core.handle.ttsReportHandle import enqueue_tts_report
|
from core.handle.ttsReportHandle import enqueue_tts_report
|
||||||
from core.providers.tts.base import audio_to_opus_data
|
from core.utils.util import audio_to_data
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
|
||||||
|
|
||||||
|
|
||||||
async def handleAudioMessage(conn, audio):
|
async def handleAudioMessage(conn, audio):
|
||||||
|
if conn.vad is None:
|
||||||
|
return
|
||||||
if not conn.asr_server_receive:
|
if not conn.asr_server_receive:
|
||||||
logger.bind(tag=TAG).debug(f"前期数据处理中,暂停接收")
|
conn.logger.bind(tag=TAG).debug(f"前期数据处理中,暂停接收")
|
||||||
return
|
return
|
||||||
if conn.client_listen_mode == "auto":
|
if conn.client_listen_mode == "auto":
|
||||||
have_voice = conn.vad.is_vad(conn, audio)
|
have_voice = conn.vad.is_vad(conn, audio)
|
||||||
@@ -40,7 +40,7 @@ async def handleAudioMessage(conn, audio):
|
|||||||
conn.asr_server_receive = True
|
conn.asr_server_receive = True
|
||||||
else:
|
else:
|
||||||
text, _ = await conn.asr.speech_to_text(conn.asr_audio, conn.session_id)
|
text, _ = await conn.asr.speech_to_text(conn.asr_audio, conn.session_id)
|
||||||
logger.bind(tag=TAG).info(f"识别文本: {text}")
|
conn.logger.bind(tag=TAG).info(f"识别文本: {text}")
|
||||||
text_len, _ = remove_punctuation_and_length(text)
|
text_len, _ = remove_punctuation_and_length(text)
|
||||||
if text_len > 0:
|
if text_len > 0:
|
||||||
# 使用自定义模块进行上报
|
# 使用自定义模块进行上报
|
||||||
@@ -111,7 +111,7 @@ async def max_out_size(conn):
|
|||||||
conn.tts_last_text_index = 0
|
conn.tts_last_text_index = 0
|
||||||
conn.llm_finish_task = True
|
conn.llm_finish_task = True
|
||||||
file_path = "config/assets/max_output_size.wav"
|
file_path = "config/assets/max_output_size.wav"
|
||||||
opus_packets, _ = audio_to_opus_data(file_path)
|
opus_packets, _ = audio_to_data(file_path)
|
||||||
conn.audio_play_queue.put((opus_packets, text, 0))
|
conn.audio_play_queue.put((opus_packets, text, 0))
|
||||||
conn.close_after_chat = True
|
conn.close_after_chat = True
|
||||||
|
|
||||||
@@ -120,7 +120,7 @@ async def check_bind_device(conn):
|
|||||||
if conn.bind_code:
|
if conn.bind_code:
|
||||||
# 确保bind_code是6位数字
|
# 确保bind_code是6位数字
|
||||||
if len(conn.bind_code) != 6:
|
if len(conn.bind_code) != 6:
|
||||||
logger.bind(tag=TAG).error(f"无效的绑定码格式: {conn.bind_code}")
|
conn.logger.bind(tag=TAG).error(f"无效的绑定码格式: {conn.bind_code}")
|
||||||
text = "绑定码格式错误,请检查配置。"
|
text = "绑定码格式错误,请检查配置。"
|
||||||
await send_stt_message(conn, text)
|
await send_stt_message(conn, text)
|
||||||
return
|
return
|
||||||
@@ -133,7 +133,7 @@ async def check_bind_device(conn):
|
|||||||
|
|
||||||
# 播放提示音
|
# 播放提示音
|
||||||
music_path = "config/assets/bind_code.wav"
|
music_path = "config/assets/bind_code.wav"
|
||||||
opus_packets, _ = audio_to_opus_data(music_path)
|
opus_packets, _ = audio_to_data(music_path)
|
||||||
conn.audio_play_queue.put((opus_packets, text, 0))
|
conn.audio_play_queue.put((opus_packets, text, 0))
|
||||||
|
|
||||||
# 逐个播放数字
|
# 逐个播放数字
|
||||||
@@ -141,10 +141,10 @@ async def check_bind_device(conn):
|
|||||||
try:
|
try:
|
||||||
digit = conn.bind_code[i]
|
digit = conn.bind_code[i]
|
||||||
num_path = f"config/assets/bind_code/{digit}.wav"
|
num_path = f"config/assets/bind_code/{digit}.wav"
|
||||||
num_packets, _ = audio_to_opus_data(num_path)
|
num_packets, _ = audio_to_data(num_path)
|
||||||
conn.audio_play_queue.put((num_packets, None, i + 1))
|
conn.audio_play_queue.put((num_packets, None, i + 1))
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"播放数字音频失败: {e}")
|
conn.logger.bind(tag=TAG).error(f"播放数字音频失败: {e}")
|
||||||
continue
|
continue
|
||||||
else:
|
else:
|
||||||
text = f"没有找到该设备的版本信息,请正确配置 OTA地址,然后重新编译固件。"
|
text = f"没有找到该设备的版本信息,请正确配置 OTA地址,然后重新编译固件。"
|
||||||
@@ -153,5 +153,5 @@ async def check_bind_device(conn):
|
|||||||
conn.tts_last_text_index = 0
|
conn.tts_last_text_index = 0
|
||||||
conn.llm_finish_task = True
|
conn.llm_finish_task = True
|
||||||
music_path = "config/assets/bind_not_found.wav"
|
music_path = "config/assets/bind_not_found.wav"
|
||||||
opus_packets, _ = audio_to_opus_data(music_path)
|
opus_packets, _ = audio_to_data(music_path)
|
||||||
conn.audio_play_queue.put((opus_packets, text, 0))
|
conn.audio_play_queue.put((opus_packets, text, 0))
|
||||||
|
|||||||
@@ -1,11 +1,9 @@
|
|||||||
from config.logger import setup_logging
|
|
||||||
import json
|
import json
|
||||||
import asyncio
|
import asyncio
|
||||||
import time
|
import time
|
||||||
from core.utils.util import get_string_no_punctuation_or_emoji, analyze_emotion
|
from core.utils.util import get_string_no_punctuation_or_emoji, analyze_emotion
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
|
||||||
|
|
||||||
emoji_map = {
|
emoji_map = {
|
||||||
"neutral": "😶",
|
"neutral": "😶",
|
||||||
@@ -49,10 +47,10 @@ async def sendAudioMessage(conn, audios, text, text_index=0):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if text_index == conn.tts_first_text_index:
|
if text_index == conn.tts_first_text_index:
|
||||||
logger.bind(tag=TAG).info(f"发送第一段语音: {text}")
|
conn.logger.bind(tag=TAG).info(f"发送第一段语音: {text}")
|
||||||
await send_tts_message(conn, "sentence_start", text)
|
await send_tts_message(conn, "sentence_start", text)
|
||||||
|
|
||||||
is_first_audio = (text_index == conn.tts_first_text_index)
|
is_first_audio = text_index == conn.tts_first_text_index
|
||||||
await sendAudio(conn, audios, pre_buffer=is_first_audio)
|
await sendAudio(conn, audios, pre_buffer=is_first_audio)
|
||||||
|
|
||||||
await send_tts_message(conn, "sentence_end", text)
|
await send_tts_message(conn, "sentence_end", text)
|
||||||
@@ -117,7 +115,7 @@ async def send_tts_message(conn, state, text=None):
|
|||||||
stop_tts_notify_voice = conn.config.get(
|
stop_tts_notify_voice = conn.config.get(
|
||||||
"stop_tts_notify_voice", "config/assets/tts_notify.mp3"
|
"stop_tts_notify_voice", "config/assets/tts_notify.mp3"
|
||||||
)
|
)
|
||||||
audios, duration = conn.tts.audio_to_opus_data(stop_tts_notify_voice)
|
audios, _ = conn.tts.audio_to_opus_data(stop_tts_notify_voice)
|
||||||
await sendAudio(conn, audios)
|
await sendAudio(conn, audios)
|
||||||
# 清除服务端讲话状态
|
# 清除服务端讲话状态
|
||||||
conn.clearSpeakStatus()
|
conn.clearSpeakStatus()
|
||||||
|
|||||||
@@ -1,4 +1,3 @@
|
|||||||
from config.logger import setup_logging
|
|
||||||
import json
|
import json
|
||||||
from core.handle.abortHandle import handleAbortMessage
|
from core.handle.abortHandle import handleAbortMessage
|
||||||
from core.handle.helloHandle import handleHelloMessage
|
from core.handle.helloHandle import handleHelloMessage
|
||||||
@@ -10,25 +9,26 @@ from core.handle.ttsReportHandle import enqueue_tts_report
|
|||||||
import asyncio
|
import asyncio
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
|
||||||
|
|
||||||
|
|
||||||
async def handleTextMessage(conn, message):
|
async def handleTextMessage(conn, message):
|
||||||
"""处理文本消息"""
|
"""处理文本消息"""
|
||||||
logger.bind(tag=TAG).info(f"收到文本消息:{message}")
|
conn.logger.bind(tag=TAG).info(f"收到文本消息:{message}")
|
||||||
try:
|
try:
|
||||||
msg_json = json.loads(message)
|
msg_json = json.loads(message)
|
||||||
if isinstance(msg_json, int):
|
if isinstance(msg_json, int):
|
||||||
await conn.websocket.send(message)
|
await conn.websocket.send(message)
|
||||||
return
|
return
|
||||||
if msg_json["type"] == "hello":
|
if msg_json["type"] == "hello":
|
||||||
await handleHelloMessage(conn)
|
await handleHelloMessage(conn, msg_json)
|
||||||
elif msg_json["type"] == "abort":
|
elif msg_json["type"] == "abort":
|
||||||
await handleAbortMessage(conn)
|
await handleAbortMessage(conn)
|
||||||
elif msg_json["type"] == "listen":
|
elif msg_json["type"] == "listen":
|
||||||
if "mode" in msg_json:
|
if "mode" in msg_json:
|
||||||
conn.client_listen_mode = msg_json["mode"]
|
conn.client_listen_mode = msg_json["mode"]
|
||||||
logger.bind(tag=TAG).debug(f"客户端拾音模式:{conn.client_listen_mode}")
|
conn.logger.bind(tag=TAG).debug(
|
||||||
|
f"客户端拾音模式:{conn.client_listen_mode}"
|
||||||
|
)
|
||||||
if msg_json["state"] == "start":
|
if msg_json["state"] == "start":
|
||||||
conn.client_have_voice = True
|
conn.client_have_voice = True
|
||||||
conn.client_voice_stop = False
|
conn.client_voice_stop = False
|
||||||
|
|||||||
@@ -9,16 +9,11 @@ TTS上报功能已集成到ConnectionHandler类中。
|
|||||||
具体实现请参考core/connection.py中的相关代码。
|
具体实现请参考core/connection.py中的相关代码。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import os
|
|
||||||
import uuid
|
|
||||||
import wave
|
|
||||||
import opuslib_next
|
import opuslib_next
|
||||||
|
|
||||||
from config.logger import setup_logging
|
|
||||||
from config.manage_api_client import report
|
from config.manage_api_client import report
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
|
||||||
|
|
||||||
|
|
||||||
def report_tts(conn, type, text, opus_data):
|
def report_tts(conn, type, text, opus_data):
|
||||||
@@ -32,7 +27,7 @@ def report_tts(conn, type, text, opus_data):
|
|||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
if opus_data:
|
if opus_data:
|
||||||
audio_data = opus_to_wav(opus_data)
|
audio_data = opus_to_wav(conn, opus_data)
|
||||||
else:
|
else:
|
||||||
audio_data = None
|
audio_data = None
|
||||||
# 执行上报
|
# 执行上报
|
||||||
@@ -44,10 +39,10 @@ def report_tts(conn, type, text, opus_data):
|
|||||||
audio=audio_data,
|
audio=audio_data,
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"TTS上报失败: {e}")
|
conn.logger.bind(tag=TAG).error(f"TTS上报失败: {e}")
|
||||||
|
|
||||||
|
|
||||||
def opus_to_wav(opus_data):
|
def opus_to_wav(conn, opus_data):
|
||||||
"""将Opus数据转换为WAV格式的字节流
|
"""将Opus数据转换为WAV格式的字节流
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -65,7 +60,7 @@ def opus_to_wav(opus_data):
|
|||||||
pcm_frame = decoder.decode(opus_packet, 960) # 960 samples = 60ms
|
pcm_frame = decoder.decode(opus_packet, 960) # 960 samples = 60ms
|
||||||
pcm_data.append(pcm_frame)
|
pcm_data.append(pcm_frame)
|
||||||
except opuslib_next.OpusError as e:
|
except opuslib_next.OpusError as e:
|
||||||
logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True)
|
conn.logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True)
|
||||||
|
|
||||||
if not pcm_data:
|
if not pcm_data:
|
||||||
raise ValueError("没有有效的PCM数据")
|
raise ValueError("没有有效的PCM数据")
|
||||||
@@ -108,8 +103,8 @@ def enqueue_tts_report(conn, type, text, opus_data):
|
|||||||
# 使用连接对象的队列,传入文本和二进制数据而非文件路径
|
# 使用连接对象的队列,传入文本和二进制数据而非文件路径
|
||||||
conn.tts_report_queue.put((type, text, opus_data))
|
conn.tts_report_queue.put((type, text, opus_data))
|
||||||
|
|
||||||
logger.bind(tag=TAG).debug(
|
conn.logger.bind(tag=TAG).debug(
|
||||||
f"TTS数据已加入上报队列: {conn.device_id}, 音频大小: {len(opus_data)} "
|
f"TTS数据已加入上报队列: {conn.device_id}, 音频大小: {len(opus_data)} "
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"加入TTS上报队列失败: {text}, {e}")
|
conn.logger.bind(tag=TAG).error(f"加入TTS上报队列失败: {text}, {e}")
|
||||||
|
|||||||
@@ -1,9 +1,9 @@
|
|||||||
"""MCP服务管理器"""
|
"""MCP服务管理器"""
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import os, json
|
import os, json
|
||||||
from typing import Dict, Any, List
|
from typing import Dict, Any, List
|
||||||
from .MCPClient import MCPClient
|
from .MCPClient import MCPClient
|
||||||
from config.logger import setup_logging
|
|
||||||
from plugins_func.register import register_function, ToolType
|
from plugins_func.register import register_function, ToolType
|
||||||
from config.config_loader import get_project_dir
|
from config.config_loader import get_project_dir
|
||||||
|
|
||||||
@@ -18,11 +18,10 @@ class MCPManager:
|
|||||||
初始化MCP管理器
|
初始化MCP管理器
|
||||||
"""
|
"""
|
||||||
self.conn = conn
|
self.conn = conn
|
||||||
self.logger = setup_logging()
|
|
||||||
self.config_path = get_project_dir() + "data/.mcp_server_settings.json"
|
self.config_path = get_project_dir() + "data/.mcp_server_settings.json"
|
||||||
if os.path.exists(self.config_path) == False:
|
if os.path.exists(self.config_path) == False:
|
||||||
self.config_path = ""
|
self.config_path = ""
|
||||||
self.logger.bind(tag=TAG).warning(
|
self.conn.logger.bind(tag=TAG).warning(
|
||||||
f"请检查mcp服务配置文件:data/.mcp_server_settings.json"
|
f"请检查mcp服务配置文件:data/.mcp_server_settings.json"
|
||||||
)
|
)
|
||||||
self.client: Dict[str, MCPClient] = {}
|
self.client: Dict[str, MCPClient] = {}
|
||||||
@@ -41,7 +40,7 @@ class MCPManager:
|
|||||||
config = json.load(f)
|
config = json.load(f)
|
||||||
return config.get("mcpServers", {})
|
return config.get("mcpServers", {})
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
self.logger.bind(tag=TAG).error(
|
self.conn.logger.bind(tag=TAG).error(
|
||||||
f"Error loading MCP config from {self.config_path}: {e}"
|
f"Error loading MCP config from {self.config_path}: {e}"
|
||||||
)
|
)
|
||||||
return {}
|
return {}
|
||||||
@@ -51,7 +50,7 @@ class MCPManager:
|
|||||||
config = self.load_config()
|
config = self.load_config()
|
||||||
for name, srv_config in config.items():
|
for name, srv_config in config.items():
|
||||||
if not srv_config.get("command") and not srv_config.get("url"):
|
if not srv_config.get("command") and not srv_config.get("url"):
|
||||||
self.logger.bind(tag=TAG).warning(
|
self.conn.logger.bind(tag=TAG).warning(
|
||||||
f"Skipping server {name}: neither command nor url specified"
|
f"Skipping server {name}: neither command nor url specified"
|
||||||
)
|
)
|
||||||
continue
|
continue
|
||||||
@@ -60,7 +59,7 @@ class MCPManager:
|
|||||||
client = MCPClient(srv_config)
|
client = MCPClient(srv_config)
|
||||||
await client.initialize()
|
await client.initialize()
|
||||||
self.client[name] = client
|
self.client[name] = client
|
||||||
self.logger.bind(tag=TAG).info(f"Initialized MCP client: {name}")
|
self.conn.logger.bind(tag=TAG).info(f"Initialized MCP client: {name}")
|
||||||
client_tools = client.get_available_tools()
|
client_tools = client.get_available_tools()
|
||||||
self.tools.extend(client_tools)
|
self.tools.extend(client_tools)
|
||||||
for tool in client_tools:
|
for tool in client_tools:
|
||||||
@@ -73,7 +72,7 @@ class MCPManager:
|
|||||||
)
|
)
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
self.logger.bind(tag=TAG).error(
|
self.conn.logger.bind(tag=TAG).error(
|
||||||
f"Failed to initialize MCP server {name}: {e}"
|
f"Failed to initialize MCP server {name}: {e}"
|
||||||
)
|
)
|
||||||
self.conn.func_handler.upload_functions_desc()
|
self.conn.func_handler.upload_functions_desc()
|
||||||
@@ -94,8 +93,8 @@ class MCPManager:
|
|||||||
"""
|
"""
|
||||||
for tool in self.tools:
|
for tool in self.tools:
|
||||||
if (
|
if (
|
||||||
tool.get("function") != None
|
tool.get("function") != None
|
||||||
and tool["function"].get("name") == tool_name
|
and tool["function"].get("name") == tool_name
|
||||||
):
|
):
|
||||||
return True
|
return True
|
||||||
return False
|
return False
|
||||||
@@ -110,7 +109,7 @@ class MCPManager:
|
|||||||
Raises:
|
Raises:
|
||||||
ValueError: 工具未找到时抛出
|
ValueError: 工具未找到时抛出
|
||||||
"""
|
"""
|
||||||
self.logger.bind(tag=TAG).info(
|
self.conn.logger.bind(tag=TAG).info(
|
||||||
f"Executing tool {tool_name} with arguments: {arguments}"
|
f"Executing tool {tool_name} with arguments: {arguments}"
|
||||||
)
|
)
|
||||||
for client in self.client.values():
|
for client in self.client.values():
|
||||||
@@ -124,9 +123,9 @@ class MCPManager:
|
|||||||
for name, client in list(self.client.items()):
|
for name, client in list(self.client.items()):
|
||||||
try:
|
try:
|
||||||
await asyncio.wait_for(client.cleanup(), timeout=20)
|
await asyncio.wait_for(client.cleanup(), timeout=20)
|
||||||
self.logger.bind(tag=TAG).info(f"MCP client closed: {name}")
|
self.conn.logger.bind(tag=TAG).info(f"MCP client closed: {name}")
|
||||||
except (asyncio.TimeoutError, Exception) as e:
|
except (asyncio.TimeoutError, Exception) as e:
|
||||||
self.logger.bind(tag=TAG).error(
|
self.conn.logger.bind(tag=TAG).error(
|
||||||
f"Error closing MCP client {name}: {e}"
|
f"Error closing MCP client {name}: {e}"
|
||||||
)
|
)
|
||||||
self.client.clear()
|
self.client.clear()
|
||||||
|
|||||||
@@ -20,64 +20,78 @@ from core.providers.asr.base import ASRProviderBase
|
|||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
|
|
||||||
|
|
||||||
class AccessToken:
|
class AccessToken:
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _encode_text(text):
|
def _encode_text(text):
|
||||||
encoded_text = parse.quote_plus(text)
|
encoded_text = parse.quote_plus(text)
|
||||||
return encoded_text.replace('+', '%20').replace('*', '%2A').replace('%7E', '~')
|
return encoded_text.replace("+", "%20").replace("*", "%2A").replace("%7E", "~")
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _encode_dict(dic):
|
def _encode_dict(dic):
|
||||||
keys = dic.keys()
|
keys = dic.keys()
|
||||||
dic_sorted = [(key, dic[key]) for key in sorted(keys)]
|
dic_sorted = [(key, dic[key]) for key in sorted(keys)]
|
||||||
encoded_text = parse.urlencode(dic_sorted)
|
encoded_text = parse.urlencode(dic_sorted)
|
||||||
return encoded_text.replace('+', '%20').replace('*', '%2A').replace('%7E', '~')
|
return encoded_text.replace("+", "%20").replace("*", "%2A").replace("%7E", "~")
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def create_token(access_key_id, access_key_secret):
|
def create_token(access_key_id, access_key_secret):
|
||||||
parameters = {'AccessKeyId': access_key_id,
|
parameters = {
|
||||||
'Action': 'CreateToken',
|
"AccessKeyId": access_key_id,
|
||||||
'Format': 'JSON',
|
"Action": "CreateToken",
|
||||||
'RegionId': 'cn-shanghai',
|
"Format": "JSON",
|
||||||
'SignatureMethod': 'HMAC-SHA1',
|
"RegionId": "cn-shanghai",
|
||||||
'SignatureNonce': str(uuid.uuid1()),
|
"SignatureMethod": "HMAC-SHA1",
|
||||||
'SignatureVersion': '1.0',
|
"SignatureNonce": str(uuid.uuid1()),
|
||||||
'Timestamp': time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()),
|
"SignatureVersion": "1.0",
|
||||||
'Version': '2019-02-28'}
|
"Timestamp": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()),
|
||||||
|
"Version": "2019-02-28",
|
||||||
|
}
|
||||||
# 构造规范化的请求字符串
|
# 构造规范化的请求字符串
|
||||||
query_string = AccessToken._encode_dict(parameters)
|
query_string = AccessToken._encode_dict(parameters)
|
||||||
# print('规范化的请求字符串: %s' % query_string)
|
# print('规范化的请求字符串: %s' % query_string)
|
||||||
# 构造待签名字符串
|
# 构造待签名字符串
|
||||||
string_to_sign = 'GET' + '&' + AccessToken._encode_text('/') + '&' + AccessToken._encode_text(query_string)
|
string_to_sign = (
|
||||||
|
"GET"
|
||||||
|
+ "&"
|
||||||
|
+ AccessToken._encode_text("/")
|
||||||
|
+ "&"
|
||||||
|
+ AccessToken._encode_text(query_string)
|
||||||
|
)
|
||||||
# print('待签名的字符串: %s' % string_to_sign)
|
# print('待签名的字符串: %s' % string_to_sign)
|
||||||
# 计算签名
|
# 计算签名
|
||||||
secreted_string = hmac.new(bytes(access_key_secret + '&', encoding='utf-8'),
|
secreted_string = hmac.new(
|
||||||
bytes(string_to_sign, encoding='utf-8'),
|
bytes(access_key_secret + "&", encoding="utf-8"),
|
||||||
hashlib.sha1).digest()
|
bytes(string_to_sign, encoding="utf-8"),
|
||||||
|
hashlib.sha1,
|
||||||
|
).digest()
|
||||||
signature = base64.b64encode(secreted_string)
|
signature = base64.b64encode(secreted_string)
|
||||||
# print('签名: %s' % signature)
|
# print('签名: %s' % signature)
|
||||||
# 进行URL编码
|
# 进行URL编码
|
||||||
signature = AccessToken._encode_text(signature)
|
signature = AccessToken._encode_text(signature)
|
||||||
# print('URL编码后的签名: %s' % signature)
|
# print('URL编码后的签名: %s' % signature)
|
||||||
# 调用服务
|
# 调用服务
|
||||||
full_url = 'http://nls-meta.cn-shanghai.aliyuncs.com/?Signature=%s&%s' % (signature, query_string)
|
full_url = "http://nls-meta.cn-shanghai.aliyuncs.com/?Signature=%s&%s" % (
|
||||||
|
signature,
|
||||||
|
query_string,
|
||||||
|
)
|
||||||
# print('url: %s' % full_url)
|
# print('url: %s' % full_url)
|
||||||
# 提交HTTP GET请求
|
# 提交HTTP GET请求
|
||||||
response = requests.get(full_url)
|
response = requests.get(full_url)
|
||||||
if response.ok:
|
if response.ok:
|
||||||
root_obj = response.json()
|
root_obj = response.json()
|
||||||
key = 'Token'
|
key = "Token"
|
||||||
if key in root_obj:
|
if key in root_obj:
|
||||||
token = root_obj[key]['Id']
|
token = root_obj[key]["Id"]
|
||||||
expire_time = root_obj[key]['ExpireTime']
|
expire_time = root_obj[key]["ExpireTime"]
|
||||||
return token, expire_time
|
return token, expire_time
|
||||||
# print(response.text)
|
# print(response.text)
|
||||||
return None, None
|
return None, None
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
class ASRProvider(ASRProviderBase):
|
class ASRProvider(ASRProviderBase):
|
||||||
def __init__(self, config: dict, delete_audio_file: bool):
|
def __init__(self, config: dict, delete_audio_file: bool):
|
||||||
|
super().__init__()
|
||||||
"""阿里云ASR初始化"""
|
"""阿里云ASR初始化"""
|
||||||
# 新增空值判断逻辑
|
# 新增空值判断逻辑
|
||||||
self.access_key_id = config.get("access_key_id")
|
self.access_key_id = config.get("access_key_id")
|
||||||
@@ -102,28 +116,23 @@ class ASRProvider(ASRProviderBase):
|
|||||||
# 确保输出目录存在
|
# 确保输出目录存在
|
||||||
os.makedirs(self.output_dir, exist_ok=True)
|
os.makedirs(self.output_dir, exist_ok=True)
|
||||||
|
|
||||||
|
|
||||||
def _refresh_token(self):
|
def _refresh_token(self):
|
||||||
"""刷新Token并记录过期时间"""
|
"""刷新Token并记录过期时间"""
|
||||||
if self.access_key_id and self.access_key_secret:
|
if self.access_key_id and self.access_key_secret:
|
||||||
self.token, expire_time_str = AccessToken.create_token(
|
self.token, expire_time_str = AccessToken.create_token(
|
||||||
self.access_key_id,
|
self.access_key_id, self.access_key_secret
|
||||||
self.access_key_secret
|
|
||||||
)
|
)
|
||||||
if not expire_time_str:
|
if not expire_time_str:
|
||||||
raise ValueError("无法获取有效的Token过期时间")
|
raise ValueError("无法获取有效的Token过期时间")
|
||||||
|
|
||||||
try:
|
try:
|
||||||
#统一转换为字符串处理
|
# 统一转换为字符串处理
|
||||||
expire_str = str(expire_time_str).strip()
|
expire_str = str(expire_time_str).strip()
|
||||||
|
|
||||||
if expire_str.isdigit():
|
if expire_str.isdigit():
|
||||||
expire_time = datetime.fromtimestamp(int(expire_str))
|
expire_time = datetime.fromtimestamp(int(expire_str))
|
||||||
else:
|
else:
|
||||||
expire_time = datetime.strptime(
|
expire_time = datetime.strptime(expire_str, "%Y-%m-%dT%H:%M:%SZ")
|
||||||
expire_str,
|
|
||||||
"%Y-%m-%dT%H:%M:%SZ"
|
|
||||||
)
|
|
||||||
self.expire_time = expire_time.timestamp() - 60
|
self.expire_time = expire_time.timestamp() - 60
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
raise ValueError(f"无效的过期时间格式: {expire_str}") from e
|
raise ValueError(f"无效的过期时间格式: {expire_str}") from e
|
||||||
@@ -145,9 +154,12 @@ class ASRProvider(ASRProviderBase):
|
|||||||
# f"过期时间 {datetime.fromtimestamp(self.expire_time)} | "
|
# f"过期时间 {datetime.fromtimestamp(self.expire_time)} | "
|
||||||
# f"剩余 {remaining:.2f}秒")
|
# f"剩余 {remaining:.2f}秒")
|
||||||
return time.time() > self.expire_time
|
return time.time() > self.expire_time
|
||||||
def generate_filename(self, extension=".wav"):
|
|
||||||
return os.path.join(self.output_file, f"tts-{__name__}{datetime.now().date()}@{uuid.uuid4().hex}{extension}")
|
|
||||||
|
|
||||||
|
def generate_filename(self, extension=".wav"):
|
||||||
|
return os.path.join(
|
||||||
|
self.output_file,
|
||||||
|
f"tts-{__name__}{datetime.now().date()}@{uuid.uuid4().hex}{extension}",
|
||||||
|
)
|
||||||
|
|
||||||
def _construct_request_url(self) -> str:
|
def _construct_request_url(self) -> str:
|
||||||
"""构造请求URL,包含参数"""
|
"""构造请求URL,包含参数"""
|
||||||
@@ -159,20 +171,6 @@ class ASRProvider(ASRProviderBase):
|
|||||||
request += "&enable_voice_detection=false"
|
request += "&enable_voice_detection=false"
|
||||||
return request
|
return request
|
||||||
|
|
||||||
def decode_opus(self, opus_data: List[bytes], session_id: str) -> List[bytes]:
|
|
||||||
"""将Opus数据解码为PCM"""
|
|
||||||
decoder = opuslib_next.Decoder(16000, 1) # 16kHz, 单声道
|
|
||||||
pcm_data = []
|
|
||||||
|
|
||||||
for opus_packet in opus_data:
|
|
||||||
try:
|
|
||||||
pcm_frame = decoder.decode(opus_packet, 960) # 960 samples = 60ms
|
|
||||||
pcm_data.append(pcm_frame)
|
|
||||||
except opuslib_next.OpusError as e:
|
|
||||||
logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True)
|
|
||||||
|
|
||||||
return pcm_data
|
|
||||||
|
|
||||||
def save_audio_to_file(self, pcm_data: List[bytes], session_id: str) -> str:
|
def save_audio_to_file(self, pcm_data: List[bytes], session_id: str) -> str:
|
||||||
"""PCM数据保存为WAV文件"""
|
"""PCM数据保存为WAV文件"""
|
||||||
module_name = __name__.split(".")[-1]
|
module_name = __name__.split(".")[-1]
|
||||||
@@ -183,7 +181,7 @@ class ASRProvider(ASRProviderBase):
|
|||||||
wf.setnchannels(1) # 单声道
|
wf.setnchannels(1) # 单声道
|
||||||
wf.setsampwidth(2) # 16-bit
|
wf.setsampwidth(2) # 16-bit
|
||||||
wf.setframerate(self.sample_rate)
|
wf.setframerate(self.sample_rate)
|
||||||
wf.writeframes(b''.join(pcm_data))
|
wf.writeframes(b"".join(pcm_data))
|
||||||
|
|
||||||
logger.bind(tag=TAG).debug(f"音频文件已保存至: {file_path}")
|
logger.bind(tag=TAG).debug(f"音频文件已保存至: {file_path}")
|
||||||
return file_path
|
return file_path
|
||||||
@@ -193,9 +191,9 @@ class ASRProvider(ASRProviderBase):
|
|||||||
try:
|
try:
|
||||||
# 设置HTTP头
|
# 设置HTTP头
|
||||||
headers = {
|
headers = {
|
||||||
'X-NLS-Token': self.token,
|
"X-NLS-Token": self.token,
|
||||||
'Content-type': 'application/octet-stream',
|
"Content-type": "application/octet-stream",
|
||||||
'Content-Length': str(len(pcm_data))
|
"Content-Length": str(len(pcm_data)),
|
||||||
}
|
}
|
||||||
|
|
||||||
# 创建连接并发送请求
|
# 创建连接并发送请求
|
||||||
@@ -203,12 +201,12 @@ class ASRProvider(ASRProviderBase):
|
|||||||
request_url = self._construct_request_url()
|
request_url = self._construct_request_url()
|
||||||
|
|
||||||
loop = asyncio.get_event_loop()
|
loop = asyncio.get_event_loop()
|
||||||
await loop.run_in_executor(None, lambda: conn.request(
|
await loop.run_in_executor(
|
||||||
method='POST',
|
None,
|
||||||
url=request_url,
|
lambda: conn.request(
|
||||||
body=pcm_data,
|
method="POST", url=request_url, body=pcm_data, headers=headers
|
||||||
headers=headers
|
),
|
||||||
))
|
)
|
||||||
|
|
||||||
# 获取响应
|
# 获取响应
|
||||||
response = await loop.run_in_executor(None, conn.getresponse)
|
response = await loop.run_in_executor(None, conn.getresponse)
|
||||||
@@ -218,10 +216,10 @@ class ASRProvider(ASRProviderBase):
|
|||||||
# 解析响应
|
# 解析响应
|
||||||
try:
|
try:
|
||||||
body_json = json.loads(body)
|
body_json = json.loads(body)
|
||||||
status = body_json.get('status')
|
status = body_json.get("status")
|
||||||
|
|
||||||
if status == 20000000:
|
if status == 20000000:
|
||||||
result = body_json.get('result', '')
|
result = body_json.get("result", "")
|
||||||
logger.bind(tag=TAG).debug(f"ASR结果: {result}")
|
logger.bind(tag=TAG).debug(f"ASR结果: {result}")
|
||||||
return result
|
return result
|
||||||
else:
|
else:
|
||||||
@@ -236,7 +234,9 @@ class ASRProvider(ASRProviderBase):
|
|||||||
logger.bind(tag=TAG).error(f"ASR请求失败: {e}", exc_info=True)
|
logger.bind(tag=TAG).error(f"ASR请求失败: {e}", exc_info=True)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
async def speech_to_text(self, opus_data: List[bytes], session_id: str) -> Tuple[Optional[str], Optional[str]]:
|
async def speech_to_text(
|
||||||
|
self, opus_data: List[bytes], session_id: str
|
||||||
|
) -> Tuple[Optional[str], Optional[str]]:
|
||||||
"""将语音数据转换为文本"""
|
"""将语音数据转换为文本"""
|
||||||
if self._is_token_expired():
|
if self._is_token_expired():
|
||||||
logger.warning("Token已过期,正在自动刷新...")
|
logger.warning("Token已过期,正在自动刷新...")
|
||||||
@@ -245,8 +245,11 @@ class ASRProvider(ASRProviderBase):
|
|||||||
file_path = None
|
file_path = None
|
||||||
try:
|
try:
|
||||||
# 解码Opus为PCM
|
# 解码Opus为PCM
|
||||||
pcm_data = self.decode_opus(opus_data, session_id)
|
if self.audio_format == "pcm":
|
||||||
combined_pcm_data = b''.join(pcm_data)
|
pcm_data = opus_data
|
||||||
|
else:
|
||||||
|
pcm_data = self.decode_opus(opus_data)
|
||||||
|
combined_pcm_data = b"".join(pcm_data)
|
||||||
|
|
||||||
# 判断是否保存为WAV文件
|
# 判断是否保存为WAV文件
|
||||||
if self.delete_audio_file:
|
if self.delete_audio_file:
|
||||||
@@ -264,4 +267,4 @@ class ASRProvider(ASRProviderBase):
|
|||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"语音识别失败: {e}", exc_info=True)
|
logger.bind(tag=TAG).error(f"语音识别失败: {e}", exc_info=True)
|
||||||
return "", file_path
|
return "", file_path
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ logger = setup_logging()
|
|||||||
|
|
||||||
class ASRProvider(ASRProviderBase):
|
class ASRProvider(ASRProviderBase):
|
||||||
def __init__(self, config: dict, delete_audio_file: bool = True):
|
def __init__(self, config: dict, delete_audio_file: bool = True):
|
||||||
|
super().__init__()
|
||||||
self.app_id = config.get("app_id")
|
self.app_id = config.get("app_id")
|
||||||
self.api_key = config.get("api_key")
|
self.api_key = config.get("api_key")
|
||||||
self.secret_key = config.get("secret_key")
|
self.secret_key = config.get("secret_key")
|
||||||
@@ -49,22 +50,6 @@ class ASRProvider(ASRProviderBase):
|
|||||||
|
|
||||||
return file_path
|
return file_path
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def decode_opus(opus_data: List[bytes]) -> bytes:
|
|
||||||
"""将Opus音频数据解码为PCM数据"""
|
|
||||||
|
|
||||||
decoder = opuslib_next.Decoder(16000, 1) # 16kHz, 单声道
|
|
||||||
pcm_data = []
|
|
||||||
|
|
||||||
for opus_packet in opus_data:
|
|
||||||
try:
|
|
||||||
pcm_frame = decoder.decode(opus_packet, 960) # 960 samples = 60ms
|
|
||||||
pcm_data.append(pcm_frame)
|
|
||||||
except opuslib_next.OpusError as e:
|
|
||||||
logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True)
|
|
||||||
|
|
||||||
return pcm_data
|
|
||||||
|
|
||||||
async def speech_to_text(
|
async def speech_to_text(
|
||||||
self, opus_data: List[bytes], session_id: str
|
self, opus_data: List[bytes], session_id: str
|
||||||
) -> Tuple[Optional[str], Optional[str]]:
|
) -> Tuple[Optional[str], Optional[str]]:
|
||||||
@@ -81,7 +66,10 @@ class ASRProvider(ASRProviderBase):
|
|||||||
return None, file_path
|
return None, file_path
|
||||||
|
|
||||||
# 将Opus音频数据解码为PCM
|
# 将Opus音频数据解码为PCM
|
||||||
pcm_data = self.decode_opus(opus_data)
|
if self.audio_format == "pcm":
|
||||||
|
pcm_data = opus_data
|
||||||
|
else:
|
||||||
|
pcm_data = self.decode_opus(opus_data)
|
||||||
combined_pcm_data = b"".join(pcm_data)
|
combined_pcm_data = b"".join(pcm_data)
|
||||||
|
|
||||||
# 判断是否保存为WAV文件
|
# 判断是否保存为WAV文件
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from typing import Optional, Tuple, List
|
from typing import Optional, Tuple, List
|
||||||
|
import opuslib_next
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
@@ -8,12 +8,37 @@ logger = setup_logging()
|
|||||||
|
|
||||||
|
|
||||||
class ASRProviderBase(ABC):
|
class ASRProviderBase(ABC):
|
||||||
|
def __init__(self):
|
||||||
|
self.audio_format = "opus"
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def save_audio_to_file(self, pcm_data: List[bytes], session_id: str) -> str:
|
def save_audio_to_file(self, pcm_data: List[bytes], session_id: str) -> str:
|
||||||
"""PCM数据保存为WAV文件"""
|
"""PCM数据保存为WAV文件"""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
async def speech_to_text(self, opus_data: List[bytes], session_id: str) -> Tuple[Optional[str], Optional[str]]:
|
async def speech_to_text(
|
||||||
|
self, opus_data: List[bytes], session_id: str
|
||||||
|
) -> Tuple[Optional[str], Optional[str]]:
|
||||||
"""将语音数据转换为文本"""
|
"""将语音数据转换为文本"""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
def set_audio_format(self, format: str) -> None:
|
||||||
|
"""设置音频格式"""
|
||||||
|
self.audio_format = format
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def decode_opus(opus_data: List[bytes]) -> bytes:
|
||||||
|
"""将Opus音频数据解码为PCM数据"""
|
||||||
|
|
||||||
|
decoder = opuslib_next.Decoder(16000, 1) # 16kHz, 单声道
|
||||||
|
pcm_data = []
|
||||||
|
|
||||||
|
for opus_packet in opus_data:
|
||||||
|
try:
|
||||||
|
pcm_frame = decoder.decode(opus_packet, 960) # 960 samples = 60ms
|
||||||
|
pcm_data.append(pcm_frame)
|
||||||
|
except opuslib_next.OpusError as e:
|
||||||
|
logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True)
|
||||||
|
|
||||||
|
return pcm_data
|
||||||
|
|||||||
@@ -85,11 +85,12 @@ def parse_response(res):
|
|||||||
|
|
||||||
class ASRProvider(ASRProviderBase):
|
class ASRProvider(ASRProviderBase):
|
||||||
def __init__(self, config: dict, delete_audio_file: bool):
|
def __init__(self, config: dict, delete_audio_file: bool):
|
||||||
|
super().__init__()
|
||||||
self.appid = config.get("appid")
|
self.appid = config.get("appid")
|
||||||
self.cluster = config.get("cluster")
|
self.cluster = config.get("cluster")
|
||||||
self.access_token = config.get("access_token")
|
self.access_token = config.get("access_token")
|
||||||
self.boosting_table_name = config.get("boosting_table_name")
|
self.boosting_table_name = config.get("boosting_table_name", "")
|
||||||
self.correct_table_name = config.get("correct_table_name")
|
self.correct_table_name = config.get("correct_table_name", "")
|
||||||
self.output_dir = config.get("output_dir")
|
self.output_dir = config.get("output_dir")
|
||||||
self.delete_audio_file = delete_audio_file
|
self.delete_audio_file = delete_audio_file
|
||||||
|
|
||||||
@@ -226,21 +227,6 @@ class ASRProvider(ASRProviderBase):
|
|||||||
logger.bind(tag=TAG).error(f"ASR request failed: {e}", exc_info=True)
|
logger.bind(tag=TAG).error(f"ASR request failed: {e}", exc_info=True)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def decode_opus(opus_data: List[bytes], session_id: str) -> List[bytes]:
|
|
||||||
|
|
||||||
decoder = opuslib_next.Decoder(16000, 1) # 16kHz, 单声道
|
|
||||||
pcm_data = []
|
|
||||||
|
|
||||||
for opus_packet in opus_data:
|
|
||||||
try:
|
|
||||||
pcm_frame = decoder.decode(opus_packet, 960) # 960 samples = 60ms
|
|
||||||
pcm_data.append(pcm_frame)
|
|
||||||
except opuslib_next.OpusError as e:
|
|
||||||
logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True)
|
|
||||||
|
|
||||||
return pcm_data
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def slice_data(data: bytes, chunk_size: int) -> (list, bool):
|
def slice_data(data: bytes, chunk_size: int) -> (list, bool):
|
||||||
"""
|
"""
|
||||||
@@ -265,7 +251,10 @@ class ASRProvider(ASRProviderBase):
|
|||||||
file_path = None
|
file_path = None
|
||||||
try:
|
try:
|
||||||
# 合并所有opus数据包
|
# 合并所有opus数据包
|
||||||
pcm_data = self.decode_opus(opus_data, session_id)
|
if self.audio_format == "pcm":
|
||||||
|
pcm_data = opus_data
|
||||||
|
else:
|
||||||
|
pcm_data = self.decode_opus(opus_data)
|
||||||
combined_pcm_data = b"".join(pcm_data)
|
combined_pcm_data = b"".join(pcm_data)
|
||||||
|
|
||||||
# 判断是否保存为WAV文件
|
# 判断是否保存为WAV文件
|
||||||
|
|||||||
@@ -6,9 +6,7 @@ import io
|
|||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
from typing import Optional, Tuple, List
|
from typing import Optional, Tuple, List
|
||||||
import uuid
|
import uuid
|
||||||
import opuslib_next
|
|
||||||
from core.providers.asr.base import ASRProviderBase
|
from core.providers.asr.base import ASRProviderBase
|
||||||
|
|
||||||
from funasr import AutoModel
|
from funasr import AutoModel
|
||||||
from funasr.utils.postprocess_utils import rich_transcription_postprocess
|
from funasr.utils.postprocess_utils import rich_transcription_postprocess
|
||||||
|
|
||||||
@@ -35,6 +33,7 @@ class CaptureOutput:
|
|||||||
|
|
||||||
class ASRProvider(ASRProviderBase):
|
class ASRProvider(ASRProviderBase):
|
||||||
def __init__(self, config: dict, delete_audio_file: bool):
|
def __init__(self, config: dict, delete_audio_file: bool):
|
||||||
|
super().__init__()
|
||||||
self.model_dir = config.get("model_dir")
|
self.model_dir = config.get("model_dir")
|
||||||
self.output_dir = config.get("output_dir") # 修正配置键名
|
self.output_dir = config.get("output_dir") # 修正配置键名
|
||||||
self.delete_audio_file = delete_audio_file
|
self.delete_audio_file = delete_audio_file
|
||||||
@@ -46,7 +45,7 @@ class ASRProvider(ASRProviderBase):
|
|||||||
model=self.model_dir,
|
model=self.model_dir,
|
||||||
vad_kwargs={"max_single_segment_time": 30000},
|
vad_kwargs={"max_single_segment_time": 30000},
|
||||||
disable_update=True,
|
disable_update=True,
|
||||||
hub="hf"
|
hub="hf",
|
||||||
# device="cuda:0", # 启用GPU加速
|
# device="cuda:0", # 启用GPU加速
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -64,27 +63,18 @@ class ASRProvider(ASRProviderBase):
|
|||||||
|
|
||||||
return file_path
|
return file_path
|
||||||
|
|
||||||
@staticmethod
|
async def speech_to_text(
|
||||||
def decode_opus(opus_data: List[bytes], session_id: str) -> List[bytes]:
|
self, opus_data: List[bytes], session_id: str
|
||||||
|
) -> Tuple[Optional[str], Optional[str]]:
|
||||||
decoder = opuslib_next.Decoder(16000, 1) # 16kHz, 单声道
|
|
||||||
pcm_data = []
|
|
||||||
|
|
||||||
for opus_packet in opus_data:
|
|
||||||
try:
|
|
||||||
pcm_frame = decoder.decode(opus_packet, 960) # 960 samples = 60ms
|
|
||||||
pcm_data.append(pcm_frame)
|
|
||||||
except opuslib_next.OpusError as e:
|
|
||||||
logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True)
|
|
||||||
|
|
||||||
return pcm_data
|
|
||||||
|
|
||||||
async def speech_to_text(self, opus_data: List[bytes], session_id: str) -> Tuple[Optional[str], Optional[str]]:
|
|
||||||
"""语音转文本主处理逻辑"""
|
"""语音转文本主处理逻辑"""
|
||||||
file_path = None
|
file_path = None
|
||||||
try:
|
try:
|
||||||
# 合并所有opus数据包
|
# 合并所有opus数据包
|
||||||
pcm_data = self.decode_opus(opus_data, session_id)
|
if self.audio_format == "pcm":
|
||||||
|
pcm_data = opus_data
|
||||||
|
else:
|
||||||
|
pcm_data = self.decode_opus(opus_data)
|
||||||
|
|
||||||
combined_pcm_data = b"".join(pcm_data)
|
combined_pcm_data = b"".join(pcm_data)
|
||||||
|
|
||||||
# 判断是否保存为WAV文件
|
# 判断是否保存为WAV文件
|
||||||
@@ -103,7 +93,9 @@ class ASRProvider(ASRProviderBase):
|
|||||||
batch_size_s=60,
|
batch_size_s=60,
|
||||||
)
|
)
|
||||||
text = rich_transcription_postprocess(result[0]["text"])
|
text = rich_transcription_postprocess(result[0]["text"])
|
||||||
logger.bind(tag=TAG).debug(f"语音识别耗时: {time.time() - start_time:.3f}s | 结果: {text}")
|
logger.bind(tag=TAG).debug(
|
||||||
|
f"语音识别耗时: {time.time() - start_time:.3f}s | 结果: {text}"
|
||||||
|
)
|
||||||
|
|
||||||
return text, file_path
|
return text, file_path
|
||||||
|
|
||||||
|
|||||||
@@ -9,22 +9,29 @@ import wave
|
|||||||
import websockets
|
import websockets
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
import asyncio
|
import asyncio
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
|
|
||||||
|
|
||||||
class ASRProvider(ASRProviderBase):
|
class ASRProvider(ASRProviderBase):
|
||||||
def __init__(self, config: dict, delete_audio_file: bool):
|
def __init__(self, config: dict, delete_audio_file: bool):
|
||||||
'''
|
"""
|
||||||
Initialize the ASRProvider with server configuration.
|
Initialize the ASRProvider with server configuration.
|
||||||
:param config: Dictionary containing 'host', 'port', and 'is_ssl'.
|
:param config: Dictionary containing 'host', 'port', and 'is_ssl'.
|
||||||
:param delete_audio_file: Boolean to indicate whether to delete audio files after processing.
|
:param delete_audio_file: Boolean to indicate whether to delete audio files after processing.
|
||||||
'''
|
"""
|
||||||
self.host = config.get('host', 'localhost')
|
super().__init__()
|
||||||
self.port = config.get('port', 10095)
|
self.host = config.get("host", "localhost")
|
||||||
self.is_ssl = config.get('is_ssl', True)
|
self.port = config.get("port", 10095)
|
||||||
|
self.is_ssl = config.get("is_ssl", True)
|
||||||
self.output_dir = config.get("output_dir")
|
self.output_dir = config.get("output_dir")
|
||||||
self.delete_audio_file = delete_audio_file
|
self.delete_audio_file = delete_audio_file
|
||||||
self.uri = f"wss://{self.host}:{self.port}" if self.is_ssl else f"ws://{self.host}:{self.port}"
|
self.uri = (
|
||||||
|
f"wss://{self.host}:{self.port}"
|
||||||
|
if self.is_ssl
|
||||||
|
else f"ws://{self.host}:{self.port}"
|
||||||
|
)
|
||||||
self.ssl_context = ssl.SSLContext() if self.is_ssl else None
|
self.ssl_context = ssl.SSLContext() if self.is_ssl else None
|
||||||
if self.ssl_context:
|
if self.ssl_context:
|
||||||
self.ssl_context.check_hostname = False
|
self.ssl_context.check_hostname = False
|
||||||
@@ -44,28 +51,11 @@ class ASRProvider(ASRProviderBase):
|
|||||||
|
|
||||||
return file_path
|
return file_path
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def decode_opus(opus_data: List[bytes]) -> bytes:
|
|
||||||
"""将Opus音频数据解码为PCM数据"""
|
|
||||||
decoder = opuslib_next.Decoder(16000, 1) # 16kHz, 单声道
|
|
||||||
pcm_data = []
|
|
||||||
|
|
||||||
for opus_packet in opus_data:
|
|
||||||
try:
|
|
||||||
pcm_frame = decoder.decode(opus_packet, 960) # 960 samples = 60ms
|
|
||||||
pcm_data.append(pcm_frame)
|
|
||||||
except opuslib_next.OpusError as e:
|
|
||||||
logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True)
|
|
||||||
|
|
||||||
return pcm_data
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
async def _receive_responses(self, ws) -> None:
|
async def _receive_responses(self, ws) -> None:
|
||||||
'''
|
"""
|
||||||
Asynchronous generator to receive messages from the WebSocket.
|
Asynchronous generator to receive messages from the WebSocket.
|
||||||
Yields each message as it is received.
|
Yields each message as it is received.
|
||||||
'''
|
"""
|
||||||
text = ""
|
text = ""
|
||||||
while True:
|
while True:
|
||||||
try:
|
try:
|
||||||
@@ -78,30 +68,35 @@ class ASRProvider(ASRProviderBase):
|
|||||||
else:
|
else:
|
||||||
text += response_data.get("text", "")
|
text += response_data.get("text", "")
|
||||||
except asyncio.TimeoutError:
|
except asyncio.TimeoutError:
|
||||||
logger.bind(tag=TAG).error("Timeout while waiting for response from WebSocket.")
|
logger.bind(tag=TAG).error(
|
||||||
|
"Timeout while waiting for response from WebSocket."
|
||||||
|
)
|
||||||
break
|
break
|
||||||
except websockets.exceptions.ConnectionClosed as e:
|
except websockets.exceptions.ConnectionClosed as e:
|
||||||
logger.bind(tag=TAG).error(f"WebSocket connection closed: {e}")
|
logger.bind(tag=TAG).error(f"WebSocket connection closed: {e}")
|
||||||
break
|
break
|
||||||
return text
|
return text
|
||||||
|
|
||||||
async def _send_data(self, ws, pcm_data: bytes, session_id: str) -> tuple:
|
async def _send_data(self, ws, pcm_data: bytes, session_id: str) -> tuple:
|
||||||
'''
|
"""
|
||||||
Internal method to handle WebSocket communication.
|
Internal method to handle WebSocket communication.
|
||||||
Reuses the persistent WebSocket connection if available.
|
Reuses the persistent WebSocket connection if available.
|
||||||
:param pcm_data: PCM audio data to send.
|
:param pcm_data: PCM audio data to send.
|
||||||
:param session_id: Unique session identifier.
|
:param session_id: Unique session identifier.
|
||||||
:return: Tuple containing recognized text and optional timestamp.
|
:return: Tuple containing recognized text and optional timestamp.
|
||||||
'''
|
"""
|
||||||
|
|
||||||
# Send initial configuration message
|
# Send initial configuration message
|
||||||
config_message = json.dumps({
|
config_message = json.dumps(
|
||||||
"mode": "offline",
|
{
|
||||||
"chunk_size": [5, 10, 5],
|
"mode": "offline",
|
||||||
"chunk_interval": 10,
|
"chunk_size": [5, 10, 5],
|
||||||
"wav_name": session_id,
|
"chunk_interval": 10,
|
||||||
"is_speaking": True,
|
"wav_name": session_id,
|
||||||
"itn": False
|
"is_speaking": True,
|
||||||
})
|
"itn": False,
|
||||||
|
}
|
||||||
|
)
|
||||||
await ws.send(config_message)
|
await ws.send(config_message)
|
||||||
logger.bind(tag=TAG).debug(f"Sent configuration message: {config_message}")
|
logger.bind(tag=TAG).debug(f"Sent configuration message: {config_message}")
|
||||||
|
|
||||||
@@ -114,16 +109,20 @@ class ASRProvider(ASRProviderBase):
|
|||||||
await ws.send(end_message)
|
await ws.send(end_message)
|
||||||
logger.bind(tag=TAG).debug(f"Sent end message: {end_message}")
|
logger.bind(tag=TAG).debug(f"Sent end message: {end_message}")
|
||||||
|
|
||||||
|
async def speech_to_text(
|
||||||
async def speech_to_text(self, opus_data: List[bytes], session_id: str) -> Tuple[Optional[str], Optional[str]]:
|
self, opus_data: List[bytes], session_id: str
|
||||||
'''
|
) -> Tuple[Optional[str], Optional[str]]:
|
||||||
|
"""
|
||||||
Convert speech data to text using FunASR.
|
Convert speech data to text using FunASR.
|
||||||
:param opus_data: List of Opus-encoded audio data chunks.
|
:param opus_data: List of Opus-encoded audio data chunks.
|
||||||
:param session_id: Unique session identifier.
|
:param session_id: Unique session identifier.
|
||||||
:return: Tuple containing recognized text and optional timestamp.
|
:return: Tuple containing recognized text and optional timestamp.
|
||||||
'''
|
"""
|
||||||
file_path = None
|
file_path = None
|
||||||
pcm_data = self.decode_opus(opus_data)
|
if self.audio_format == "pcm":
|
||||||
|
pcm_data = opus_data
|
||||||
|
else:
|
||||||
|
pcm_data = self.decode_opus(opus_data)
|
||||||
combined_pcm_data = b"".join(pcm_data)
|
combined_pcm_data = b"".join(pcm_data)
|
||||||
|
|
||||||
# 判断是否保存为WAV文件
|
# 判断是否保存为WAV文件
|
||||||
@@ -132,16 +131,19 @@ class ASRProvider(ASRProviderBase):
|
|||||||
else:
|
else:
|
||||||
file_path = self.save_audio_to_file(pcm_data, session_id)
|
file_path = self.save_audio_to_file(pcm_data, session_id)
|
||||||
|
|
||||||
async with websockets.connect(self.uri, subprotocols=["binary"], ping_interval=None, ssl=self.ssl_context) as ws:
|
async with websockets.connect(
|
||||||
|
self.uri, subprotocols=["binary"], ping_interval=None, ssl=self.ssl_context
|
||||||
|
) as ws:
|
||||||
try:
|
try:
|
||||||
# Use asyncio to handle WebSocket communication
|
# Use asyncio to handle WebSocket communication
|
||||||
send_task = asyncio.create_task(self._send_data(ws, combined_pcm_data, session_id))
|
send_task = asyncio.create_task(
|
||||||
|
self._send_data(ws, combined_pcm_data, session_id)
|
||||||
|
)
|
||||||
receive_task = asyncio.create_task(self._receive_responses(ws))
|
receive_task = asyncio.create_task(self._receive_responses(ws))
|
||||||
|
|
||||||
# Gather tasks with error handling
|
# Gather tasks with error handling
|
||||||
done, pending = await asyncio.wait(
|
done, pending = await asyncio.wait(
|
||||||
[send_task, receive_task],
|
[send_task, receive_task], return_when=asyncio.FIRST_EXCEPTION
|
||||||
return_when=asyncio.FIRST_EXCEPTION
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# Cancel any pending tasks
|
# Cancel any pending tasks
|
||||||
@@ -155,11 +157,16 @@ class ASRProvider(ASRProviderBase):
|
|||||||
|
|
||||||
# Get the result from the receive task
|
# Get the result from the receive task
|
||||||
result = receive_task.result()
|
result = receive_task.result()
|
||||||
return result, file_path # Return the recognized text and timestamp (if any)
|
return (
|
||||||
|
result,
|
||||||
|
file_path,
|
||||||
|
) # Return the recognized text and timestamp (if any)
|
||||||
|
|
||||||
except websockets.exceptions.ConnectionClosed as e:
|
except websockets.exceptions.ConnectionClosed as e:
|
||||||
logger.bind(tag=TAG).error(f"WebSocket connection closed: {e}")
|
logger.bind(tag=TAG).error(f"WebSocket connection closed: {e}")
|
||||||
return "", file_path
|
return "", file_path
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"Error during speech-to-text conversion: {e}", exc_info=True)
|
logger.bind(tag=TAG).error(
|
||||||
return "", file_path
|
f"Error during speech-to-text conversion: {e}", exc_info=True
|
||||||
|
)
|
||||||
|
return "", file_path
|
||||||
|
|||||||
@@ -37,6 +37,7 @@ class CaptureOutput:
|
|||||||
|
|
||||||
class ASRProvider(ASRProviderBase):
|
class ASRProvider(ASRProviderBase):
|
||||||
def __init__(self, config: dict, delete_audio_file: bool):
|
def __init__(self, config: dict, delete_audio_file: bool):
|
||||||
|
super().__init__()
|
||||||
self.model_dir = config.get("model_dir")
|
self.model_dir = config.get("model_dir")
|
||||||
self.output_dir = config.get("output_dir")
|
self.output_dir = config.get("output_dir")
|
||||||
self.delete_audio_file = delete_audio_file
|
self.delete_audio_file = delete_audio_file
|
||||||
@@ -97,21 +98,6 @@ class ASRProvider(ASRProviderBase):
|
|||||||
|
|
||||||
return file_path
|
return file_path
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def decode_opus(opus_data: List[bytes], session_id: str) -> List[bytes]:
|
|
||||||
|
|
||||||
decoder = opuslib_next.Decoder(16000, 1) # 16kHz, 单声道
|
|
||||||
pcm_data = []
|
|
||||||
|
|
||||||
for opus_packet in opus_data:
|
|
||||||
try:
|
|
||||||
pcm_frame = decoder.decode(opus_packet, 960) # 960 samples = 60ms
|
|
||||||
pcm_data.append(pcm_frame)
|
|
||||||
except opuslib_next.OpusError as e:
|
|
||||||
logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True)
|
|
||||||
|
|
||||||
return pcm_data
|
|
||||||
|
|
||||||
def read_wave(self, wave_filename: str) -> Tuple[np.ndarray, int]:
|
def read_wave(self, wave_filename: str) -> Tuple[np.ndarray, int]:
|
||||||
"""
|
"""
|
||||||
Args:
|
Args:
|
||||||
@@ -144,7 +130,10 @@ class ASRProvider(ASRProviderBase):
|
|||||||
try:
|
try:
|
||||||
# 保存音频文件
|
# 保存音频文件
|
||||||
start_time = time.time()
|
start_time = time.time()
|
||||||
pcm_data = self.decode_opus(opus_data, session_id)
|
if self.audio_format == "pcm":
|
||||||
|
pcm_data = opus_data
|
||||||
|
else:
|
||||||
|
pcm_data = self.decode_opus(opus_data)
|
||||||
file_path = self.save_audio_to_file(pcm_data, session_id)
|
file_path = self.save_audio_to_file(pcm_data, session_id)
|
||||||
logger.bind(tag=TAG).debug(
|
logger.bind(tag=TAG).debug(
|
||||||
f"音频文件保存耗时: {time.time() - start_time:.3f}s | 路径: {file_path}"
|
f"音频文件保存耗时: {time.time() - start_time:.3f}s | 路径: {file_path}"
|
||||||
|
|||||||
@@ -17,12 +17,14 @@ from config.logger import setup_logging
|
|||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
|
|
||||||
|
|
||||||
class ASRProvider(ASRProviderBase):
|
class ASRProvider(ASRProviderBase):
|
||||||
API_URL = "https://asr.tencentcloudapi.com"
|
API_URL = "https://asr.tencentcloudapi.com"
|
||||||
API_VERSION = "2019-06-14"
|
API_VERSION = "2019-06-14"
|
||||||
FORMAT = "pcm" # 支持的音频格式:pcm, wav, mp3
|
FORMAT = "pcm" # 支持的音频格式:pcm, wav, mp3
|
||||||
|
|
||||||
def __init__(self, config: dict, delete_audio_file: bool = True):
|
def __init__(self, config: dict, delete_audio_file: bool = True):
|
||||||
|
super().__init__()
|
||||||
self.secret_id = config.get("secret_id")
|
self.secret_id = config.get("secret_id")
|
||||||
self.secret_key = config.get("secret_key")
|
self.secret_key = config.get("secret_key")
|
||||||
self.output_dir = config.get("output_dir")
|
self.output_dir = config.get("output_dir")
|
||||||
@@ -45,23 +47,9 @@ class ASRProvider(ASRProviderBase):
|
|||||||
|
|
||||||
return file_path
|
return file_path
|
||||||
|
|
||||||
@staticmethod
|
async def speech_to_text(
|
||||||
def decode_opus(opus_data: List[bytes]) -> bytes:
|
self, opus_data: List[bytes], session_id: str
|
||||||
"""将Opus音频数据解码为PCM数据"""
|
) -> Tuple[Optional[str], Optional[str]]:
|
||||||
|
|
||||||
decoder = opuslib_next.Decoder(16000, 1) # 16kHz, 单声道
|
|
||||||
pcm_data = []
|
|
||||||
|
|
||||||
for opus_packet in opus_data:
|
|
||||||
try:
|
|
||||||
pcm_frame = decoder.decode(opus_packet, 960) # 960 samples = 60ms
|
|
||||||
pcm_data.append(pcm_frame)
|
|
||||||
except opuslib_next.OpusError as e:
|
|
||||||
logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True)
|
|
||||||
|
|
||||||
return pcm_data
|
|
||||||
|
|
||||||
async def speech_to_text(self, opus_data: List[bytes], session_id: str) -> Tuple[Optional[str], Optional[str]]:
|
|
||||||
"""将语音数据转换为文本"""
|
"""将语音数据转换为文本"""
|
||||||
if not opus_data:
|
if not opus_data:
|
||||||
logger.bind(tag=TAG).warn("音频数据为空!")
|
logger.bind(tag=TAG).warn("音频数据为空!")
|
||||||
@@ -75,7 +63,10 @@ class ASRProvider(ASRProviderBase):
|
|||||||
return None, file_path
|
return None, file_path
|
||||||
|
|
||||||
# 将Opus音频数据解码为PCM
|
# 将Opus音频数据解码为PCM
|
||||||
pcm_data = self.decode_opus(opus_data)
|
if self.audio_format == "pcm":
|
||||||
|
pcm_data = opus_data
|
||||||
|
else:
|
||||||
|
pcm_data = self.decode_opus(opus_data)
|
||||||
combined_pcm_data = b"".join(pcm_data)
|
combined_pcm_data = b"".join(pcm_data)
|
||||||
|
|
||||||
# 判断是否保存为WAV文件
|
# 判断是否保存为WAV文件
|
||||||
@@ -85,7 +76,7 @@ class ASRProvider(ASRProviderBase):
|
|||||||
self.save_audio_to_file(pcm_data, session_id)
|
self.save_audio_to_file(pcm_data, session_id)
|
||||||
|
|
||||||
# 将音频数据转换为Base64编码
|
# 将音频数据转换为Base64编码
|
||||||
base64_audio = base64.b64encode(combined_pcm_data).decode('utf-8')
|
base64_audio = base64.b64encode(combined_pcm_data).decode("utf-8")
|
||||||
|
|
||||||
# 构建请求体
|
# 构建请求体
|
||||||
request_body = self._build_request_body(base64_audio)
|
request_body = self._build_request_body(base64_audio)
|
||||||
@@ -98,7 +89,9 @@ class ASRProvider(ASRProviderBase):
|
|||||||
result = self._send_request(request_body, timestamp, authorization)
|
result = self._send_request(request_body, timestamp, authorization)
|
||||||
|
|
||||||
if result:
|
if result:
|
||||||
logger.bind(tag=TAG).debug(f"腾讯云语音识别耗时: {time.time() - start_time:.3f}s | 结果: {result}")
|
logger.bind(tag=TAG).debug(
|
||||||
|
f"腾讯云语音识别耗时: {time.time() - start_time:.3f}s | 结果: {result}"
|
||||||
|
)
|
||||||
|
|
||||||
return result, file_path
|
return result, file_path
|
||||||
|
|
||||||
@@ -115,7 +108,7 @@ class ASRProvider(ASRProviderBase):
|
|||||||
"SourceType": 1, # 音频数据来源为语音文件
|
"SourceType": 1, # 音频数据来源为语音文件
|
||||||
"VoiceFormat": self.FORMAT, # 音频格式
|
"VoiceFormat": self.FORMAT, # 音频格式
|
||||||
"Data": base64_audio, # Base64编码的音频数据
|
"Data": base64_audio, # Base64编码的音频数据
|
||||||
"DataLen": len(base64_audio) # 数据长度
|
"DataLen": len(base64_audio), # 数据长度
|
||||||
}
|
}
|
||||||
return json.dumps(request_map)
|
return json.dumps(request_map)
|
||||||
|
|
||||||
@@ -148,9 +141,11 @@ class ASRProvider(ASRProviderBase):
|
|||||||
action = "SentenceRecognition" # 接口名称
|
action = "SentenceRecognition" # 接口名称
|
||||||
|
|
||||||
# 构建规范头部信息,注意顺序和格式
|
# 构建规范头部信息,注意顺序和格式
|
||||||
canonical_headers = f"content-type:{content_type.lower()}\n" + \
|
canonical_headers = (
|
||||||
f"host:{host.lower()}\n" + \
|
f"content-type:{content_type.lower()}\n"
|
||||||
f"x-tc-action:{action.lower()}\n"
|
+ f"host:{host.lower()}\n"
|
||||||
|
+ f"x-tc-action:{action.lower()}\n"
|
||||||
|
)
|
||||||
|
|
||||||
signed_headers = "content-type;host;x-tc-action"
|
signed_headers = "content-type;host;x-tc-action"
|
||||||
|
|
||||||
@@ -158,21 +153,25 @@ class ASRProvider(ASRProviderBase):
|
|||||||
payload_hash = self._sha256_hex(request_body)
|
payload_hash = self._sha256_hex(request_body)
|
||||||
|
|
||||||
# 构建规范请求字符串
|
# 构建规范请求字符串
|
||||||
canonical_request = f"{http_request_method}\n" + \
|
canonical_request = (
|
||||||
f"{canonical_uri}\n" + \
|
f"{http_request_method}\n"
|
||||||
f"{canonical_query_string}\n" + \
|
+ f"{canonical_uri}\n"
|
||||||
f"{canonical_headers}\n" + \
|
+ f"{canonical_query_string}\n"
|
||||||
f"{signed_headers}\n" + \
|
+ f"{canonical_headers}\n"
|
||||||
f"{payload_hash}"
|
+ f"{signed_headers}\n"
|
||||||
|
+ f"{payload_hash}"
|
||||||
|
)
|
||||||
|
|
||||||
# 计算规范请求的哈希值
|
# 计算规范请求的哈希值
|
||||||
hashed_canonical_request = self._sha256_hex(canonical_request)
|
hashed_canonical_request = self._sha256_hex(canonical_request)
|
||||||
|
|
||||||
# 构建待签名字符串
|
# 构建待签名字符串
|
||||||
string_to_sign = f"{algorithm}\n" + \
|
string_to_sign = (
|
||||||
f"{timestamp}\n" + \
|
f"{algorithm}\n"
|
||||||
f"{credential_scope}\n" + \
|
+ f"{timestamp}\n"
|
||||||
f"{hashed_canonical_request}"
|
+ f"{credential_scope}\n"
|
||||||
|
+ f"{hashed_canonical_request}"
|
||||||
|
)
|
||||||
|
|
||||||
# 计算签名密钥
|
# 计算签名密钥
|
||||||
secret_date = self._hmac_sha256(f"TC3{self.secret_key}", date)
|
secret_date = self._hmac_sha256(f"TC3{self.secret_key}", date)
|
||||||
@@ -180,13 +179,17 @@ class ASRProvider(ASRProviderBase):
|
|||||||
secret_signing = self._hmac_sha256(secret_service, "tc3_request")
|
secret_signing = self._hmac_sha256(secret_service, "tc3_request")
|
||||||
|
|
||||||
# 计算签名
|
# 计算签名
|
||||||
signature = self._bytes_to_hex(self._hmac_sha256(secret_signing, string_to_sign))
|
signature = self._bytes_to_hex(
|
||||||
|
self._hmac_sha256(secret_signing, string_to_sign)
|
||||||
|
)
|
||||||
|
|
||||||
# 构建授权头
|
# 构建授权头
|
||||||
authorization = f"{algorithm} " + \
|
authorization = (
|
||||||
f"Credential={self.secret_id}/{credential_scope}, " + \
|
f"{algorithm} "
|
||||||
f"SignedHeaders={signed_headers}, " + \
|
+ f"Credential={self.secret_id}/{credential_scope}, "
|
||||||
f"Signature={signature}"
|
+ f"SignedHeaders={signed_headers}, "
|
||||||
|
+ f"Signature={signature}"
|
||||||
|
)
|
||||||
|
|
||||||
return timestamp, authorization
|
return timestamp, authorization
|
||||||
|
|
||||||
@@ -194,7 +197,9 @@ class ASRProvider(ASRProviderBase):
|
|||||||
logger.bind(tag=TAG).error(f"生成认证头失败: {e}", exc_info=True)
|
logger.bind(tag=TAG).error(f"生成认证头失败: {e}", exc_info=True)
|
||||||
raise RuntimeError(f"生成认证头失败: {e}")
|
raise RuntimeError(f"生成认证头失败: {e}")
|
||||||
|
|
||||||
def _send_request(self, request_body: str, timestamp: str, authorization: str) -> Optional[str]:
|
def _send_request(
|
||||||
|
self, request_body: str, timestamp: str, authorization: str
|
||||||
|
) -> Optional[str]:
|
||||||
"""发送请求到腾讯云API"""
|
"""发送请求到腾讯云API"""
|
||||||
headers = {
|
headers = {
|
||||||
"Content-Type": "application/json; charset=utf-8",
|
"Content-Type": "application/json; charset=utf-8",
|
||||||
@@ -203,7 +208,7 @@ class ASRProvider(ASRProviderBase):
|
|||||||
"X-TC-Action": "SentenceRecognition",
|
"X-TC-Action": "SentenceRecognition",
|
||||||
"X-TC-Version": self.API_VERSION,
|
"X-TC-Version": self.API_VERSION,
|
||||||
"X-TC-Timestamp": timestamp,
|
"X-TC-Timestamp": timestamp,
|
||||||
"X-TC-Region": "ap-shanghai"
|
"X-TC-Region": "ap-shanghai",
|
||||||
}
|
}
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -234,16 +239,16 @@ class ASRProvider(ASRProviderBase):
|
|||||||
|
|
||||||
def _sha256_hex(self, data: str) -> str:
|
def _sha256_hex(self, data: str) -> str:
|
||||||
"""计算字符串的SHA256哈希值"""
|
"""计算字符串的SHA256哈希值"""
|
||||||
digest = hashlib.sha256(data.encode('utf-8')).digest()
|
digest = hashlib.sha256(data.encode("utf-8")).digest()
|
||||||
return self._bytes_to_hex(digest)
|
return self._bytes_to_hex(digest)
|
||||||
|
|
||||||
def _hmac_sha256(self, key, data: str) -> bytes:
|
def _hmac_sha256(self, key, data: str) -> bytes:
|
||||||
"""计算HMAC-SHA256"""
|
"""计算HMAC-SHA256"""
|
||||||
if isinstance(key, str):
|
if isinstance(key, str):
|
||||||
key = key.encode('utf-8')
|
key = key.encode("utf-8")
|
||||||
|
|
||||||
return hmac.new(key, data.encode('utf-8'), hashlib.sha256).digest()
|
return hmac.new(key, data.encode("utf-8"), hashlib.sha256).digest()
|
||||||
|
|
||||||
def _bytes_to_hex(self, bytes_data: bytes) -> str:
|
def _bytes_to_hex(self, bytes_data: bytes) -> str:
|
||||||
"""字节数组转十六进制字符串"""
|
"""字节数组转十六进制字符串"""
|
||||||
return ''.join(f"{b:02x}" for b in bytes_data)
|
return "".join(f"{b:02x}" for b in bytes_data)
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ from config.logger import setup_logging
|
|||||||
import os
|
import os
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from core.utils.tts import MarkdownCleaner
|
from core.utils.tts import MarkdownCleaner
|
||||||
from core.utils.util import audio_to_opus_data
|
from core.utils.util import audio_to_data
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
@@ -53,5 +53,10 @@ class TTSProviderBase(ABC):
|
|||||||
async def text_to_speak(self, text, output_file):
|
async def text_to_speak(self, text, output_file):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
def audio_to_pcm_data(self, audio_file_path):
|
||||||
|
"""音频文件转换为PCM编码"""
|
||||||
|
return audio_to_data(audio_file_path, is_opus=False)
|
||||||
|
|
||||||
def audio_to_opus_data(self, audio_file_path):
|
def audio_to_opus_data(self, audio_file_path):
|
||||||
return audio_to_opus_data(audio_file_path)
|
"""音频文件转换为Opus编码"""
|
||||||
|
return audio_to_data(audio_file_path, is_opus=True)
|
||||||
|
|||||||
@@ -4,7 +4,14 @@ from datetime import datetime
|
|||||||
|
|
||||||
|
|
||||||
class Message:
|
class Message:
|
||||||
def __init__(self, role: str, content: str = None, uniq_id: str = None, tool_calls = None, tool_call_id=None):
|
def __init__(
|
||||||
|
self,
|
||||||
|
role: str,
|
||||||
|
content: str = None,
|
||||||
|
uniq_id: str = None,
|
||||||
|
tool_calls=None,
|
||||||
|
tool_call_id=None,
|
||||||
|
):
|
||||||
self.uniq_id = uniq_id if uniq_id is not None else str(uuid.uuid4())
|
self.uniq_id = uniq_id if uniq_id is not None else str(uuid.uuid4())
|
||||||
self.role = role
|
self.role = role
|
||||||
self.content = content
|
self.content = content
|
||||||
@@ -16,7 +23,7 @@ class Dialogue:
|
|||||||
def __init__(self):
|
def __init__(self):
|
||||||
self.dialogue: List[Message] = []
|
self.dialogue: List[Message] = []
|
||||||
# 获取当前时间
|
# 获取当前时间
|
||||||
self.current_time = datetime.now().strftime('%Y-%m-%d %H:%M:%S')
|
self.current_time = datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||||||
|
|
||||||
def put(self, message: Message):
|
def put(self, message: Message):
|
||||||
self.dialogue.append(message)
|
self.dialogue.append(message)
|
||||||
@@ -25,7 +32,9 @@ class Dialogue:
|
|||||||
if m.tool_calls is not None:
|
if m.tool_calls is not None:
|
||||||
dialogue.append({"role": m.role, "tool_calls": m.tool_calls})
|
dialogue.append({"role": m.role, "tool_calls": m.tool_calls})
|
||||||
elif m.role == "tool":
|
elif m.role == "tool":
|
||||||
dialogue.append({"role": m.role, "tool_call_id": m.tool_call_id, "content": m.content})
|
dialogue.append(
|
||||||
|
{"role": m.role, "tool_call_id": m.tool_call_id, "content": m.content}
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
dialogue.append({"role": m.role, "content": m.content})
|
dialogue.append({"role": m.role, "content": m.content})
|
||||||
|
|
||||||
@@ -44,23 +53,23 @@ class Dialogue:
|
|||||||
else:
|
else:
|
||||||
self.put(Message(role="system", content=new_content))
|
self.put(Message(role="system", content=new_content))
|
||||||
|
|
||||||
def get_llm_dialogue_with_memory(self, memory_str: str = None) -> List[Dict[str, str]]:
|
def get_llm_dialogue_with_memory(
|
||||||
|
self, memory_str: str = None
|
||||||
|
) -> List[Dict[str, str]]:
|
||||||
if memory_str is None or len(memory_str) == 0:
|
if memory_str is None or len(memory_str) == 0:
|
||||||
return self.get_llm_dialogue()
|
return self.get_llm_dialogue()
|
||||||
|
|
||||||
# 构建带记忆的对话
|
# 构建带记忆的对话
|
||||||
dialogue = []
|
dialogue = []
|
||||||
|
|
||||||
# 添加系统提示和记忆
|
# 添加系统提示和记忆
|
||||||
system_message = next(
|
system_message = next(
|
||||||
(msg for msg in self.dialogue if msg.role == "system"), None
|
(msg for msg in self.dialogue if msg.role == "system"), None
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
if system_message:
|
if system_message:
|
||||||
enhanced_system_prompt = (
|
enhanced_system_prompt = (
|
||||||
f"{system_message.content}\n\n"
|
f"{system_message.content}\n\n" f"相关记忆:\n{memory_str}"
|
||||||
f"相关记忆:\n{memory_str}"
|
|
||||||
)
|
)
|
||||||
dialogue.append({"role": "system", "content": enhanced_system_prompt})
|
dialogue.append({"role": "system", "content": enhanced_system_prompt})
|
||||||
|
|
||||||
|
|||||||
@@ -350,12 +350,6 @@ def initialize_modules(
|
|||||||
str(config.get("delete_audio", True)).lower() in ("true", "1", "yes"),
|
str(config.get("delete_audio", True)).lower() in ("true", "1", "yes"),
|
||||||
)
|
)
|
||||||
logger.bind(tag=TAG).info(f"初始化组件: asr成功 {select_asr_module}")
|
logger.bind(tag=TAG).info(f"初始化组件: asr成功 {select_asr_module}")
|
||||||
|
|
||||||
# 初始化自定义prompt
|
|
||||||
if config.get("prompt", None) is not None:
|
|
||||||
modules["prompt"] = config["prompt"]
|
|
||||||
logger.bind(tag=TAG).info(f"初始化组件: prompt成功 {modules['prompt'][:50]}...")
|
|
||||||
|
|
||||||
return modules
|
return modules
|
||||||
|
|
||||||
|
|
||||||
@@ -868,8 +862,7 @@ def analyze_emotion(text):
|
|||||||
return top_emotions[0] # 如果都不在优先级列表里,返回第一个
|
return top_emotions[0] # 如果都不在优先级列表里,返回第一个
|
||||||
|
|
||||||
|
|
||||||
def audio_to_opus_data(audio_file_path):
|
def audio_to_data(audio_file_path, is_opus=True):
|
||||||
"""音频文件转换为Opus编码"""
|
|
||||||
# 获取文件后缀名
|
# 获取文件后缀名
|
||||||
file_type = os.path.splitext(audio_file_path)[1]
|
file_type = os.path.splitext(audio_file_path)[1]
|
||||||
if file_type:
|
if file_type:
|
||||||
@@ -895,7 +888,7 @@ def audio_to_opus_data(audio_file_path):
|
|||||||
frame_duration = 60 # 60ms per frame
|
frame_duration = 60 # 60ms per frame
|
||||||
frame_size = int(16000 * frame_duration / 1000) # 960 samples/frame
|
frame_size = int(16000 * frame_duration / 1000) # 960 samples/frame
|
||||||
|
|
||||||
opus_datas = []
|
datas = []
|
||||||
# 按帧处理所有音频数据(包括最后一帧可能补零)
|
# 按帧处理所有音频数据(包括最后一帧可能补零)
|
||||||
for i in range(0, len(raw_data), frame_size * 2): # 16bit=2bytes/sample
|
for i in range(0, len(raw_data), frame_size * 2): # 16bit=2bytes/sample
|
||||||
# 获取当前帧的二进制数据
|
# 获取当前帧的二进制数据
|
||||||
@@ -905,11 +898,62 @@ def audio_to_opus_data(audio_file_path):
|
|||||||
if len(chunk) < frame_size * 2:
|
if len(chunk) < frame_size * 2:
|
||||||
chunk += b"\x00" * (frame_size * 2 - len(chunk))
|
chunk += b"\x00" * (frame_size * 2 - len(chunk))
|
||||||
|
|
||||||
# 转换为numpy数组处理
|
if is_opus:
|
||||||
np_frame = np.frombuffer(chunk, dtype=np.int16)
|
# 转换为numpy数组处理
|
||||||
|
np_frame = np.frombuffer(chunk, dtype=np.int16)
|
||||||
|
# 编码Opus数据
|
||||||
|
frame_data = encoder.encode(np_frame.tobytes(), frame_size)
|
||||||
|
else:
|
||||||
|
frame_data = chunk if isinstance(chunk, bytes) else bytes(chunk)
|
||||||
|
|
||||||
# 编码Opus数据
|
datas.append(frame_data)
|
||||||
opus_data = encoder.encode(np_frame.tobytes(), frame_size)
|
|
||||||
opus_datas.append(opus_data)
|
|
||||||
|
|
||||||
return opus_datas, duration
|
return datas, duration
|
||||||
|
|
||||||
|
|
||||||
|
def check_vad_update(before_config, new_config):
|
||||||
|
if (
|
||||||
|
new_config.get("selected_module") is None
|
||||||
|
or new_config["selected_module"].get("VAD") is None
|
||||||
|
):
|
||||||
|
return False
|
||||||
|
update_vad = False
|
||||||
|
current_vad_module = before_config["selected_module"]["VAD"]
|
||||||
|
new_vad_module = new_config["selected_module"]["VAD"]
|
||||||
|
current_vad_type = (
|
||||||
|
current_vad_module
|
||||||
|
if "type" not in before_config["VAD"][current_vad_module]
|
||||||
|
else before_config["VAD"][current_vad_module]["type"]
|
||||||
|
)
|
||||||
|
new_vad_type = (
|
||||||
|
new_vad_module
|
||||||
|
if "type" not in new_config["VAD"][new_vad_module]
|
||||||
|
else new_config["VAD"][new_vad_module]["type"]
|
||||||
|
)
|
||||||
|
print(f"前vad:{current_vad_type},后vad:{new_vad_type}")
|
||||||
|
update_vad = current_vad_type != new_vad_type
|
||||||
|
return update_vad
|
||||||
|
|
||||||
|
|
||||||
|
def check_asr_update(before_config, new_config):
|
||||||
|
if (
|
||||||
|
new_config.get("selected_module") is None
|
||||||
|
or new_config["selected_module"].get("ASR") is None
|
||||||
|
):
|
||||||
|
return False
|
||||||
|
update_asr = False
|
||||||
|
current_asr_module = before_config["selected_module"]["ASR"]
|
||||||
|
new_asr_module = new_config["selected_module"]["ASR"]
|
||||||
|
current_asr_type = (
|
||||||
|
current_asr_module
|
||||||
|
if "type" not in before_config["ASR"][current_asr_module]
|
||||||
|
else before_config["ASR"][current_asr_module]["type"]
|
||||||
|
)
|
||||||
|
new_asr_type = (
|
||||||
|
new_asr_module
|
||||||
|
if "type" not in new_config["ASR"][new_asr_module]
|
||||||
|
else new_config["ASR"][new_asr_module]["type"]
|
||||||
|
)
|
||||||
|
print(f"前asr:{current_asr_type},后asr:{new_asr_type}")
|
||||||
|
update_asr = current_asr_type != new_asr_type
|
||||||
|
return update_asr
|
||||||
|
|||||||
@@ -2,7 +2,7 @@ import asyncio
|
|||||||
import websockets
|
import websockets
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
from core.connection import ConnectionHandler
|
from core.connection import ConnectionHandler
|
||||||
from core.utils.util import initialize_modules
|
from core.utils.util import initialize_modules, check_vad_update, check_asr_update
|
||||||
from config.config_loader import get_config_from_api
|
from config.config_loader import get_config_from_api
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
@@ -84,38 +84,8 @@ class WebSocketServer:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
# 检查 VAD 和 ASR 类型是否需要更新
|
# 检查 VAD 和 ASR 类型是否需要更新
|
||||||
update_vad = False
|
update_vad = check_vad_update(self.config, new_config)
|
||||||
update_asr = False
|
update_asr = check_asr_update(self.config, new_config)
|
||||||
|
|
||||||
# 获取当前和新的 VAD 类型
|
|
||||||
current_vad_module = self.config["selected_module"]["VAD"]
|
|
||||||
new_vad_module = new_config["selected_module"]["VAD"]
|
|
||||||
current_vad_type = (
|
|
||||||
current_vad_module
|
|
||||||
if "type" not in self.config["VAD"][current_vad_module]
|
|
||||||
else self.config["VAD"][current_vad_module]["type"]
|
|
||||||
)
|
|
||||||
new_vad_type = (
|
|
||||||
new_vad_module
|
|
||||||
if "type" not in new_config["VAD"][new_vad_module]
|
|
||||||
else new_config["VAD"][new_vad_module]["type"]
|
|
||||||
)
|
|
||||||
update_vad = current_vad_type != new_vad_type
|
|
||||||
|
|
||||||
# 获取当前和新的 ASR 类型
|
|
||||||
current_asr_module = self.config["selected_module"]["ASR"]
|
|
||||||
new_asr_module = new_config["selected_module"]["ASR"]
|
|
||||||
current_asr_type = (
|
|
||||||
current_asr_module
|
|
||||||
if "type" not in self.config["ASR"][current_asr_module]
|
|
||||||
else self.config["ASR"][current_asr_module]["type"]
|
|
||||||
)
|
|
||||||
new_asr_type = (
|
|
||||||
new_asr_module
|
|
||||||
if "type" not in new_config["ASR"][new_asr_module]
|
|
||||||
else new_config["ASR"][new_asr_module]["type"]
|
|
||||||
)
|
|
||||||
update_asr = current_asr_type != new_asr_type
|
|
||||||
|
|
||||||
# 更新配置
|
# 更新配置
|
||||||
self.config = new_config
|
self.config = new_config
|
||||||
@@ -132,12 +102,18 @@ class WebSocketServer:
|
|||||||
)
|
)
|
||||||
|
|
||||||
# 更新组件实例
|
# 更新组件实例
|
||||||
self._vad = modules["vad"] if "vad" in modules else None
|
if "vad" in modules:
|
||||||
self._asr = modules["asr"] if "asr" in modules else None
|
self._vad = modules["vad"]
|
||||||
self._tts = modules["tts"] if "tts" in modules else None
|
if "asr" in modules:
|
||||||
self._llm = modules["llm"] if "llm" in modules else None
|
self._asr = modules["asr"]
|
||||||
self._intent = modules["intent"] if "intent" in modules else None
|
if "tts" in modules:
|
||||||
self._memory = modules["memory"] if "memory" in modules else None
|
self._tts = modules["tts"]
|
||||||
|
if "llm" in modules:
|
||||||
|
self._llm = modules["llm"]
|
||||||
|
if "intent" in modules:
|
||||||
|
self._intent = modules["intent"]
|
||||||
|
if "memory" in modules:
|
||||||
|
self._memory = modules["memory"]
|
||||||
|
|
||||||
return True
|
return True
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
|||||||
@@ -10,10 +10,9 @@ from pathlib import Path
|
|||||||
from core.utils import p3
|
from core.utils import p3
|
||||||
from core.handle.sendAudioHandle import send_stt_message
|
from core.handle.sendAudioHandle import send_stt_message
|
||||||
from plugins_func.register import register_function, ToolType, ActionResponse, Action
|
from plugins_func.register import register_function, ToolType, ActionResponse, Action
|
||||||
|
from core.utils.dialogue import Message
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
|
||||||
|
|
||||||
MUSIC_CACHE = {}
|
MUSIC_CACHE = {}
|
||||||
|
|
||||||
@@ -45,7 +44,7 @@ def play_music(conn, song_name: str):
|
|||||||
|
|
||||||
# 检查事件循环状态
|
# 检查事件循环状态
|
||||||
if not conn.loop.is_running():
|
if not conn.loop.is_running():
|
||||||
logger.bind(tag=TAG).error("事件循环未运行,无法提交任务")
|
conn.logger.bind(tag=TAG).error("事件循环未运行,无法提交任务")
|
||||||
return ActionResponse(
|
return ActionResponse(
|
||||||
action=Action.RESPONSE, result="系统繁忙", response="请稍后再试"
|
action=Action.RESPONSE, result="系统繁忙", response="请稍后再试"
|
||||||
)
|
)
|
||||||
@@ -59,9 +58,9 @@ def play_music(conn, song_name: str):
|
|||||||
def handle_done(f):
|
def handle_done(f):
|
||||||
try:
|
try:
|
||||||
f.result() # 可在此处理成功逻辑
|
f.result() # 可在此处理成功逻辑
|
||||||
logger.bind(tag=TAG).info("播放完成")
|
conn.logger.bind(tag=TAG).info("播放完成")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"播放失败: {e}")
|
conn.logger.bind(tag=TAG).error(f"播放失败: {e}")
|
||||||
|
|
||||||
future.add_done_callback(handle_done)
|
future.add_done_callback(handle_done)
|
||||||
|
|
||||||
@@ -69,7 +68,7 @@ def play_music(conn, song_name: str):
|
|||||||
action=Action.NONE, result="指令已接收", response="正在为您播放音乐"
|
action=Action.NONE, result="指令已接收", response="正在为您播放音乐"
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"处理音乐意图错误: {e}")
|
conn.logger.bind(tag=TAG).error(f"处理音乐意图错误: {e}")
|
||||||
return ActionResponse(
|
return ActionResponse(
|
||||||
action=Action.RESPONSE, result=str(e), response="播放音乐时出错了"
|
action=Action.RESPONSE, result=str(e), response="播放音乐时出错了"
|
||||||
)
|
)
|
||||||
@@ -150,7 +149,7 @@ async def handle_music_command(conn, text):
|
|||||||
|
|
||||||
"""处理音乐播放指令"""
|
"""处理音乐播放指令"""
|
||||||
clean_text = re.sub(r"[^\w\s]", "", text).strip()
|
clean_text = re.sub(r"[^\w\s]", "", text).strip()
|
||||||
logger.bind(tag=TAG).debug(f"检查是否是音乐命令: {clean_text}")
|
conn.logger.bind(tag=TAG).debug(f"检查是否是音乐命令: {clean_text}")
|
||||||
|
|
||||||
# 尝试匹配具体歌名
|
# 尝试匹配具体歌名
|
||||||
if os.path.exists(MUSIC_CACHE["music_dir"]):
|
if os.path.exists(MUSIC_CACHE["music_dir"]):
|
||||||
@@ -165,7 +164,7 @@ async def handle_music_command(conn, text):
|
|||||||
if potential_song:
|
if potential_song:
|
||||||
best_match = _find_best_match(potential_song, MUSIC_CACHE["music_files"])
|
best_match = _find_best_match(potential_song, MUSIC_CACHE["music_files"])
|
||||||
if best_match:
|
if best_match:
|
||||||
logger.bind(tag=TAG).info(f"找到最匹配的歌曲: {best_match}")
|
conn.logger.bind(tag=TAG).info(f"找到最匹配的歌曲: {best_match}")
|
||||||
await play_local_music(conn, specific_file=best_match)
|
await play_local_music(conn, specific_file=best_match)
|
||||||
return True
|
return True
|
||||||
# 检查是否是通用播放音乐命令
|
# 检查是否是通用播放音乐命令
|
||||||
@@ -195,7 +194,9 @@ async def play_local_music(conn, specific_file=None):
|
|||||||
"""播放本地音乐文件"""
|
"""播放本地音乐文件"""
|
||||||
try:
|
try:
|
||||||
if not os.path.exists(MUSIC_CACHE["music_dir"]):
|
if not os.path.exists(MUSIC_CACHE["music_dir"]):
|
||||||
logger.bind(tag=TAG).error(f"音乐目录不存在: " + MUSIC_CACHE["music_dir"])
|
conn.logger.bind(tag=TAG).error(
|
||||||
|
f"音乐目录不存在: " + MUSIC_CACHE["music_dir"]
|
||||||
|
)
|
||||||
return
|
return
|
||||||
|
|
||||||
# 确保路径正确性
|
# 确保路径正确性
|
||||||
@@ -204,16 +205,17 @@ async def play_local_music(conn, specific_file=None):
|
|||||||
music_path = os.path.join(MUSIC_CACHE["music_dir"], specific_file)
|
music_path = os.path.join(MUSIC_CACHE["music_dir"], specific_file)
|
||||||
else:
|
else:
|
||||||
if not MUSIC_CACHE["music_files"]:
|
if not MUSIC_CACHE["music_files"]:
|
||||||
logger.bind(tag=TAG).error("未找到MP3音乐文件")
|
conn.logger.bind(tag=TAG).error("未找到MP3音乐文件")
|
||||||
return
|
return
|
||||||
selected_music = random.choice(MUSIC_CACHE["music_files"])
|
selected_music = random.choice(MUSIC_CACHE["music_files"])
|
||||||
music_path = os.path.join(MUSIC_CACHE["music_dir"], selected_music)
|
music_path = os.path.join(MUSIC_CACHE["music_dir"], selected_music)
|
||||||
|
|
||||||
if not os.path.exists(music_path):
|
if not os.path.exists(music_path):
|
||||||
logger.bind(tag=TAG).error(f"选定的音乐文件不存在: {music_path}")
|
conn.logger.bind(tag=TAG).error(f"选定的音乐文件不存在: {music_path}")
|
||||||
return
|
return
|
||||||
text = _get_random_play_prompt(selected_music)
|
text = _get_random_play_prompt(selected_music)
|
||||||
await send_stt_message(conn, text)
|
await send_stt_message(conn, text)
|
||||||
|
conn.dialogue.put(Message(role="assistant", content=text))
|
||||||
conn.tts_first_text_index = 0
|
conn.tts_first_text_index = 0
|
||||||
conn.tts_last_text_index = 0
|
conn.tts_last_text_index = 0
|
||||||
|
|
||||||
@@ -233,5 +235,5 @@ async def play_local_music(conn, specific_file=None):
|
|||||||
conn.audio_play_queue.put((opus_packets, None, conn.tts_last_text_index))
|
conn.audio_play_queue.put((opus_packets, None, conn.tts_last_text_index))
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"播放音乐失败: {str(e)}")
|
conn.logger.bind(tag=TAG).error(f"播放音乐失败: {str(e)}")
|
||||||
logger.bind(tag=TAG).error(f"详细错误: {traceback.format_exc()}")
|
conn.logger.bind(tag=TAG).error(f"详细错误: {traceback.format_exc()}")
|
||||||
|
|||||||
@@ -1,51 +1,56 @@
|
|||||||
from plugins_func.register import register_function,ToolType, ActionResponse, Action
|
from plugins_func.register import register_function, ToolType, ActionResponse, Action
|
||||||
from config.logger import setup_logging
|
|
||||||
|
|
||||||
TAG = __name__
|
|
||||||
logger = setup_logging()
|
|
||||||
|
|
||||||
plugin_loader_function_desc = {
|
plugin_loader_function_desc = {
|
||||||
"type": "function",
|
"type": "function",
|
||||||
"function": {
|
"function": {
|
||||||
"name": "plugin_loader",
|
"name": "plugin_loader",
|
||||||
"description": "当用户想加载或卸载插件/function时,调用此函数:支持的插件列表为[plugins]",
|
"description": "当用户想加载或卸载插件/function时,调用此函数:支持的插件列表为[plugins]",
|
||||||
"parameters": {
|
"parameters": {
|
||||||
"type": "object",
|
"type": "object",
|
||||||
"properties": {
|
"properties": {
|
||||||
"oper": {
|
"oper": {"type": "string", "description": "load or unload"},
|
||||||
"type": "string",
|
"name": {"type": "string", "description": "要加载或卸载的插件名字"},
|
||||||
"description": "load or unload"
|
},
|
||||||
},
|
"required": ["oper", "name"],
|
||||||
"name":{
|
},
|
||||||
"type": "string",
|
},
|
||||||
"description": "要加载或卸载的插件名字"
|
}
|
||||||
}
|
|
||||||
},
|
|
||||||
"required": ["oper","name"]
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
@register_function('plugin_loader', plugin_loader_function_desc, ToolType.SYSTEM_CTL)
|
|
||||||
|
@register_function("plugin_loader", plugin_loader_function_desc, ToolType.SYSTEM_CTL)
|
||||||
def plugin_loader(conn, oper: str, name: str):
|
def plugin_loader(conn, oper: str, name: str):
|
||||||
"""插件加载"""
|
"""插件加载"""
|
||||||
if oper not in ["load", "unload"]:
|
if oper not in ["load", "unload"]:
|
||||||
return ActionResponse(action=Action.RESPONSE, result="插件操作失败", response="不支持的操作")
|
return ActionResponse(
|
||||||
|
action=Action.RESPONSE, result="插件操作失败", response="不支持的操作"
|
||||||
|
)
|
||||||
|
|
||||||
cur_support = conn.func_handler.current_support_functions()
|
cur_support = conn.func_handler.current_support_functions()
|
||||||
if oper == "load":
|
if oper == "load":
|
||||||
if name in cur_support:
|
if name in cur_support:
|
||||||
return ActionResponse(action=Action.RESPONSE, result="插件加载失败", response=f"{name}插件已加载,无需重复加载")
|
return ActionResponse(
|
||||||
|
action=Action.RESPONSE,
|
||||||
|
result="插件加载失败",
|
||||||
|
response=f"{name}插件已加载,无需重复加载",
|
||||||
|
)
|
||||||
func = conn.func_handler.function_registry.register_function(name)
|
func = conn.func_handler.function_registry.register_function(name)
|
||||||
if not func:
|
if not func:
|
||||||
return ActionResponse(action=Action.RESPONSE, result="插件加载失败", response="插件未找到")
|
return ActionResponse(
|
||||||
|
action=Action.RESPONSE, result="插件加载失败", response="插件未找到"
|
||||||
|
)
|
||||||
res = f"{name}插件加载成功"
|
res = f"{name}插件加载成功"
|
||||||
else:
|
else:
|
||||||
if name not in cur_support:
|
if name not in cur_support:
|
||||||
return ActionResponse(action=Action.RESPONSE, result="插件卸载失败", response=f"{name}插件未加载")
|
return ActionResponse(
|
||||||
|
action=Action.RESPONSE,
|
||||||
|
result="插件卸载失败",
|
||||||
|
response=f"{name}插件未加载",
|
||||||
|
)
|
||||||
bOK = conn.func_handler.function_registry.unregister_function(name)
|
bOK = conn.func_handler.function_registry.unregister_function(name)
|
||||||
if not bOK:
|
if not bOK:
|
||||||
return ActionResponse(action=Action.RESPONSE, result="插件卸载失败", response="插件未找到")
|
return ActionResponse(
|
||||||
|
action=Action.RESPONSE, result="插件卸载失败", response="插件未找到"
|
||||||
|
)
|
||||||
res = f"{name}插件卸载成功"
|
res = f"{name}插件卸载成功"
|
||||||
conn.func_handler.upload_functions_desc()
|
conn.func_handler.upload_functions_desc()
|
||||||
return ActionResponse(action=Action.RESPONSE, result="插件操作成功", response=res)
|
return ActionResponse(action=Action.RESPONSE, result="插件操作成功", response=res)
|
||||||
|
|||||||
Reference in New Issue
Block a user