diff --git a/README.md b/README.md index c3e0a3f4..9eee3b89 100644 --- a/README.md +++ b/README.md @@ -6,7 +6,7 @@ 本项目为开源智能硬件项目 xiaozhi-esp32提供后端服务
根据小智通信协议使用Python、Java、Vue实现
-帮助您快速搭建小智服务器 +支持MCP接入点和声纹识别

@@ -230,11 +230,12 @@ Websocket接口地址: wss://2662r3426b.vicp.fun/xiaozhi/v1/ | 功能模块 | 描述 | |:---:|:---| -| 核心服务架构 | 基于WebSocket和HTTP服务器,提供完整的控制台管理和认证系统 | -| 语音交互系统 | 支持流式ASR(语音识别)、流式TTS(语音合成)、VAD(语音活动检测),支持多语言识别和语音处理 | -| 智能对话系统 | 支持多种LLM(大语言模型),实现智能对话 | -| 视觉感知系统 | 支持多种VLLM(视觉大模型),实现多模态交互 | -| 意图识别系统 | 支持LLM意图识别、Function Call函数调用,提供插件化意图处理机制 | +| 核心架构 | 基于WebSocket和HTTP服务器,提供完整的控制台管理和认证系统 | +| 语音交互 | 支持流式ASR(语音识别)、流式TTS(语音合成)、VAD(语音活动检测),支持多语言识别和语音处理 | +| 声纹识别 | 支持多用户声纹注册、管理和识别,与ASR并行处理,实时识别说话人身份并传递给LLM进行个性化回应 | +| 智能对话 | 支持多种LLM(大语言模型),实现智能对话 | +| 视觉感知 | 支持多种VLLM(视觉大模型),实现多模态交互 | +| 意图识别 | 支持LLM意图识别、Function Call函数调用,提供插件化意图处理机制 | | 记忆系统 | 支持本地短期记忆、mem0ai接口记忆,具备记忆总结功能 | | 工具调用 | 支持客户端IOT协议、客户MCP协议、服务端MCP协议、MCP接入点协议、自定义工具函数 | | 管理后台 | 提供Web管理界面,支持用户管理、系统配置和设备管理 | @@ -313,6 +314,14 @@ Websocket接口地址: wss://2662r3426b.vicp.fun/xiaozhi/v1/ --- +### Voiceprint 声纹识别 + +| 使用方式 | 支持平台 | 免费平台 | +|:---:|:---:|:---:| +| 本地使用 | 3D-Speaker | 3D-Speaker | + +--- + ### Memory 记忆存储 | 类型 | 平台名称 | 使用方式 | 收费模式 | 备注 | diff --git a/docs/FAQ.md b/docs/FAQ.md index e2d428e3..8a884d7d 100644 --- a/docs/FAQ.md +++ b/docs/FAQ.md @@ -112,7 +112,17 @@ VAD: 参考教程[视觉模型使用指南](./mcp-vision-integration.md) -### 10、更多问题,可联系我们反馈 💬 +### 10、如何开启MCP接入点 🔧 + +1、先参考教程[MCP 接入点部署使用指南](./mcp-endpoint-enable.md) + +2、再参考教程[MCP 接入点使用指南](./mcp-endpoint-integration.md) + +### 12、如何开启声纹识别 🔊 + +参考教程[声纹识别启用指南](./voiceprint-integration.md) + +### 13、更多问题,可联系我们反馈 💬 可以在[issues](https://github.com/xinnan-tech/xiaozhi-esp32-server/issues)提交您的问题。 diff --git a/docs/mcp-endpoint-enable.md b/docs/mcp-endpoint-enable.md index 7bc72484..6b64ccac 100644 --- a/docs/mcp-endpoint-enable.md +++ b/docs/mcp-endpoint-enable.md @@ -1,9 +1,9 @@ # MCP 接入点部署使用指南 -本教程包含2个部分 +本教程包含3个部分 - 1、如何部署MCP接入点这个服务 -- 1、全模块部署时,怎么配置MCP接入点 -- 2、单模块部署时,怎么配置MCP接入点 +- 2、全模块部署时,怎么配置MCP接入点 +- 3、单模块部署时,怎么配置MCP接入点 # 1、如何部署MCP接入点这个服务 diff --git a/docs/voiceprint-integration.md b/docs/voiceprint-integration.md new file mode 100644 index 00000000..c961dbb7 --- /dev/null +++ b/docs/voiceprint-integration.md @@ -0,0 +1,96 @@ +# 声纹识别启用指南 + +本教程包含2个部分 +- 1、如何部署声纹识别这个服务 +- 2、全模块部署时,怎么配置声纹识别接口 + +# 1、如何部署声纹识别这个服务 + +## 第一步,下载声纹识别项目源码 + +浏览器打开[声纹识别项目地址](https://github.com/xinnan-tech/voiceprint-api) + +打开完,找到页面中一个绿色的按钮,写着`Code`的按钮,点开它,然后你就看到`Download ZIP`的按钮。 + +点击它,下载本项目源码压缩包。下载到你电脑后,解压它,此时它的名字可能叫`voiceprint-api-main` +你需要把它重命名成`voiceprint-api`。 + +## 第二步,启动程序 +这个项目是一个很简单的项目,建议使用docker运行。不过如果你不想使用docker运行,你可以参考[这个页面](https://github.com/xinnan-tech/voiceprint-api/blob/main/README.md)使用源码运行。以下是docker运行的方法 + +``` +# 进入本项目源码根目录 +cd voiceprint-api + +# 清除缓存 +docker compose -f docker-compose.yml down +docker stop voiceprint-api +docker rm voiceprint-api +docker rmi ghcr.nju.edu.cn/xinnan-tech/voiceprint-api:latest + +# 启动docker容器 +docker compose -f docker-compose.yml up -d +# 查看日志 +docker logs -f voiceprint-api +``` + +此时,日志里会输出类似以下的日志 +``` +250711 INFO-🚀 开始: 生产环境服务启动(Uvicorn),监听地址: 0.0.0.0:8005 +250711 INFO-============================================================ +250711 INFO-声纹接口地址: http://127.0.0.1:8005/voiceprint/health?key=abcd +250711 INFO-============================================================ +``` + +请你把声纹接口地址复制出来: + +由于你是docker部署,切不可直接使用上面的地址! + +由于你是docker部署,切不可直接使用上面的地址! + +由于你是docker部署,切不可直接使用上面的地址! + +你先把地址复制出来,放在一个草稿里,你要知道你的电脑的局域网ip是什么,例如我的电脑局域网ip是`192.168.1.25`,那么 +原来我的接口地址 +``` +http://127.0.0.1:8005/voiceprint/health?key=abcd + +``` +就要改成 +``` +http://192.168.1.25:8005/voiceprint/health?key=abcd +``` + +改好后,请使用浏览器直接访问`声纹接口地址`。当浏览器出现类似这样的代码,说明是成功了。 +``` +{"total_voiceprints":0,"status":"healthy"} +``` + +请你保留好修改后的`声纹接口地址`,下一步要用到。 + +# 2、全模块部署时,怎么配置声纹识别 + +## 第一步 配置接口 +如果你是全模块部署,使用管理员账号,登录智控台,点击顶部`参数字典`,选择`参数管理`功能。 + +然后搜索参数`server.voice_print`,此时,它的值应该是`null`值。 +点击修改按钮,把上一步得来的`声纹接口地址`粘贴到`参数值`里。然后保存。 + +如果能保存成功,说明一切顺利,你可以去智能体查看效果了。如果不成功,说明智控台无法访问声纹识别,很大概率是网络防火墙,或者没有填写正确的局域网ip。 + +## 第二步 设置智能体记忆模式 + +进入你的智能体的角色配置里,将记忆设置成`本地短期记忆`,一定要开启`上报文字+语音`。 + +## 第三步 和你的智能体聊天 + +将你的设备通电,然后和他用正常的语速和音调聊天。 + +## 第四步 设置声纹 + +在智控台,`智能体管理`页面,在智能体的面板里,有一个`声纹识别`按钮,点击它。在底部有一个`新增按钮`。就可以对某个人说的话进行声纹注册。 +在弹出的框里,`描述`这个属性建议填写上,可以是这个人的职业、性格、爱好。方便智能体对说话人进行分析和了解。 + +## 第三步 和你的智能体聊天 + +将你的设备通电,问它,你知道我是谁吗?如果他能回答得出,说明声纹识别功能正常。 \ No newline at end of file diff --git a/main/manager-api/src/main/java/xiaozhi/common/constant/Constant.java b/main/manager-api/src/main/java/xiaozhi/common/constant/Constant.java index 5be7b639..011e99e9 100644 --- a/main/manager-api/src/main/java/xiaozhi/common/constant/Constant.java +++ b/main/manager-api/src/main/java/xiaozhi/common/constant/Constant.java @@ -116,6 +116,11 @@ public interface Constant { */ String SERVER_MCP_ENDPOINT = "server.mcp_endpoint"; + /** + * mcp接入点路径 + */ + String SERVER_VOICE_PRINT = "server.voice_print"; + /** * 无记忆 */ @@ -232,7 +237,7 @@ public interface Constant { /** * 版本号 */ - public static final String VERSION = "0.6.3"; + public static final String VERSION = "0.7.1"; /** * 无效固件URL diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/Enums/AgentChatHistoryType.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/Enums/AgentChatHistoryType.java new file mode 100644 index 00000000..49ed2dce --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/Enums/AgentChatHistoryType.java @@ -0,0 +1,21 @@ +package xiaozhi.modules.agent.Enums; + + +import lombok.Getter; + +/** + * 智能体聊天记录类型 + */ +@Getter +public enum AgentChatHistoryType { + + USER((byte) 1), + AGENT((byte) 2); + + private final byte value; + + AgentChatHistoryType(byte i) { + this.value = i; + } + +} diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/controller/AgentController.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/controller/AgentController.java index bfd82e41..a4ad0863 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/agent/controller/AgentController.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/controller/AgentController.java @@ -47,6 +47,7 @@ import xiaozhi.modules.agent.service.AgentChatHistoryService; import xiaozhi.modules.agent.service.AgentPluginMappingService; import xiaozhi.modules.agent.service.AgentService; import xiaozhi.modules.agent.service.AgentTemplateService; +import xiaozhi.modules.agent.vo.AgentChatHistoryUserVO; import xiaozhi.modules.agent.vo.AgentInfoVO; import xiaozhi.modules.device.entity.DeviceEntity; import xiaozhi.modules.device.service.DeviceService; @@ -181,6 +182,33 @@ public class AgentController { List result = agentChatHistoryService.getChatHistoryBySessionId(id, sessionId); return new Result>().ok(result); } + @GetMapping("/{id}/chat-history/user") + @Operation(summary = "获取智能体聊天记录(用户)") + @RequiresPermissions("sys:role:normal") + public Result> getRecentlyFiftyByAgentId( + @PathVariable("id") String id) { + // 获取当前用户 + UserDetail user = SecurityUser.getUser(); + + // 检查权限 + if (!agentService.checkAgentPermission(id, user.getId())) { + return new Result>().error("没有权限查看该智能体的聊天记录"); + } + + // 查询聊天记录 + List data = agentChatHistoryService.getRecentlyFiftyByAgentId(id); + return new Result>().ok(data); + } + + @GetMapping("/{id}/chat-history/audio") + @Operation(summary = "获取音频内容") + @RequiresPermissions("sys:role:normal") + public Result getContentByAudioId( + @PathVariable("id") String id) { + // 查询聊天记录 + String data = agentChatHistoryService.getContentByAudioId(id); + return new Result().ok(data); + } @PostMapping("/audio/{audioId}") @Operation(summary = "获取音频下载ID") diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/controller/AgentVoicePrintController.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/controller/AgentVoicePrintController.java new file mode 100644 index 00000000..1615cee5 --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/controller/AgentVoicePrintController.java @@ -0,0 +1,86 @@ +package xiaozhi.modules.agent.controller; + +import java.util.List; + +import org.apache.commons.lang3.StringUtils; +import org.apache.shiro.authz.annotation.RequiresPermissions; +import org.springframework.web.bind.annotation.DeleteMapping; +import org.springframework.web.bind.annotation.GetMapping; +import org.springframework.web.bind.annotation.PathVariable; +import org.springframework.web.bind.annotation.PostMapping; +import org.springframework.web.bind.annotation.PutMapping; +import org.springframework.web.bind.annotation.RequestBody; +import org.springframework.web.bind.annotation.RequestMapping; +import org.springframework.web.bind.annotation.RestController; + +import io.swagger.v3.oas.annotations.Operation; +import io.swagger.v3.oas.annotations.tags.Tag; +import jakarta.validation.Valid; +import lombok.AllArgsConstructor; +import xiaozhi.common.exception.RenException; +import xiaozhi.common.utils.Result; +import xiaozhi.modules.agent.dto.AgentVoicePrintSaveDTO; +import xiaozhi.modules.agent.dto.AgentVoicePrintUpdateDTO; +import xiaozhi.modules.agent.service.AgentVoicePrintService; +import xiaozhi.modules.agent.vo.AgentVoicePrintVO; +import xiaozhi.modules.security.user.SecurityUser; +import xiaozhi.modules.sys.service.SysParamsService; + +@Tag(name = "智能体声纹管理") +@AllArgsConstructor +@RestController +@RequestMapping("/agent/voice-print") +public class AgentVoicePrintController { + private final AgentVoicePrintService agentVoicePrintService; + private final SysParamsService sysParamsService; + + @PostMapping + @Operation(summary = "创建智能体的声纹") + @RequiresPermissions("sys:role:normal") + public Result save(@RequestBody @Valid AgentVoicePrintSaveDTO dto) { + boolean b = agentVoicePrintService.insert(dto); + if (b) { + return new Result<>(); + } + return new Result().error("智能体的声纹创建失败"); + } + + @PutMapping + @Operation(summary = "更新智能体的对应声纹") + @RequiresPermissions("sys:role:normal") + public Result update(@RequestBody @Valid AgentVoicePrintUpdateDTO dto) { + Long userId = SecurityUser.getUserId(); + boolean b = agentVoicePrintService.update(userId, dto); + if (b) { + return new Result<>(); + } + return new Result().error("智能体的对应声纹更新失败"); + } + + @DeleteMapping("/{id}") + @Operation(summary = "删除智能体对应声纹") + @RequiresPermissions("sys:role:normal") + public Result delete(@PathVariable String id) { + Long userId = SecurityUser.getUserId(); + // 先删除关联的设备 + boolean delete = agentVoicePrintService.delete(userId, id); + if (delete) { + return new Result<>(); + } + return new Result().error("智能体的对应声纹删除失败"); + } + + @GetMapping("/list/{id}") + @Operation(summary = "获取用户指定智能体声纹列表") + @RequiresPermissions("sys:role:normal") + public Result> list(@PathVariable String id) { + String voiceprintUrl = sysParamsService.getValue("server.voice_print", true); + if (StringUtils.isBlank(voiceprintUrl) || "null".equals(voiceprintUrl)) { + throw new RenException("声纹接口未配置,请先在参数配置中配置声纹接口地址(server.voice_print)"); + } + Long userId = SecurityUser.getUserId(); + List list = agentVoicePrintService.list(userId, id); + return new Result>().ok(list); + } + +} diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/dao/AgentVoicePrintDao.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/dao/AgentVoicePrintDao.java new file mode 100644 index 00000000..1a5ee364 --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/dao/AgentVoicePrintDao.java @@ -0,0 +1,20 @@ +package xiaozhi.modules.agent.dao; + +import org.apache.ibatis.annotations.Mapper; + +import com.baomidou.mybatisplus.core.mapper.BaseMapper; + +import xiaozhi.modules.agent.entity.AgentChatHistoryEntity; +import xiaozhi.modules.agent.entity.AgentVoicePrintEntity; + +/** + * {@link AgentChatHistoryEntity} 智能体聊天历史记录Dao对象 + * + * @author Goody + * @version 1.0, 2025/4/30 + * @since 1.0.0 + */ +@Mapper +public interface AgentVoicePrintDao extends BaseMapper { + +} diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/dto/AgentVoicePrintSaveDTO.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/dto/AgentVoicePrintSaveDTO.java new file mode 100644 index 00000000..0ff3b0ca --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/dto/AgentVoicePrintSaveDTO.java @@ -0,0 +1,28 @@ +package xiaozhi.modules.agent.dto; + +import lombok.Data; + +/** + * 保存智能体声纹的dto + * + * @author zjy + */ +@Data +public class AgentVoicePrintSaveDTO { + /** + * 关联的智能体id + */ + private String agentId; + /** + * 音频文件id + */ + private String audioId; + /** + * 声纹来源的人姓名 + */ + private String sourceName; + /** + * 描述声纹来源的人 + */ + private String introduce; +} diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/dto/AgentVoicePrintUpdateDTO.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/dto/AgentVoicePrintUpdateDTO.java new file mode 100644 index 00000000..60643e29 --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/dto/AgentVoicePrintUpdateDTO.java @@ -0,0 +1,28 @@ +package xiaozhi.modules.agent.dto; + +import lombok.Data; + +/** + * 修改智能体声纹的dto + * + * @author zjy + */ +@Data +public class AgentVoicePrintUpdateDTO { + /** + * 智能体声纹id + */ + private String id; + /** + * 音频文件id + */ + private String audioId; + /** + * 声纹来源的人姓名 + */ + private String sourceName; + /** + * 描述声纹来源的人 + */ + private String introduce; +} diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/dto/IdentifyVoicePrintResponse.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/dto/IdentifyVoicePrintResponse.java new file mode 100644 index 00000000..14ce57c2 --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/dto/IdentifyVoicePrintResponse.java @@ -0,0 +1,21 @@ +package xiaozhi.modules.agent.dto; + +import com.fasterxml.jackson.annotation.JsonProperty; + +import lombok.Data; + +/** + * 声纹识别接口返回的对象 + */ +@Data +public class IdentifyVoicePrintResponse { + /** + * 最匹配的声纹id + */ + @JsonProperty("speaker_id") + private String speakerId; + /** + * 声纹的分数 + */ + private Double score; +} diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/entity/AgentVoicePrintEntity.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/entity/AgentVoicePrintEntity.java new file mode 100644 index 00000000..4998b288 --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/entity/AgentVoicePrintEntity.java @@ -0,0 +1,64 @@ +package xiaozhi.modules.agent.entity; + +import java.util.Date; + +import com.baomidou.mybatisplus.annotation.FieldFill; +import com.baomidou.mybatisplus.annotation.IdType; +import com.baomidou.mybatisplus.annotation.TableField; +import com.baomidou.mybatisplus.annotation.TableId; +import com.baomidou.mybatisplus.annotation.TableName; + +import lombok.Data; + +/** + * 智能体声纹表 + * + * @author zjy + */ +@TableName(value = "ai_agent_voice_print") +@Data +public class AgentVoicePrintEntity { + /** + * 主键id + */ + @TableId(type = IdType.ASSIGN_UUID) + private String id; + /** + * 关联的智能体id + */ + private String agentId; + /** + * 关联的音频id + */ + private String audioId; + /** + * 声纹来源的人姓名 + */ + private String sourceName; + /** + * 描述声纹来源的人 + */ + private String introduce; + + /** + * 创建者 + */ + @TableField(fill = FieldFill.INSERT) + private Long creator; + /** + * 创建时间 + */ + @TableField(fill = FieldFill.INSERT) + private Date createDate; + + /** + * 更新者 + */ + @TableField(fill = FieldFill.INSERT_UPDATE) + private Long updater; + /** + * 更新时间 + */ + @TableField(fill = FieldFill.INSERT_UPDATE) + private Date updateDate; +} diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/service/AgentChatHistoryService.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/AgentChatHistoryService.java index 21fce988..459b5a5a 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/agent/service/AgentChatHistoryService.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/AgentChatHistoryService.java @@ -9,6 +9,7 @@ import xiaozhi.common.page.PageData; import xiaozhi.modules.agent.dto.AgentChatHistoryDTO; import xiaozhi.modules.agent.dto.AgentChatSessionDTO; import xiaozhi.modules.agent.entity.AgentChatHistoryEntity; +import xiaozhi.modules.agent.vo.AgentChatHistoryUserVO; /** * 智能体聊天记录表处理service @@ -44,4 +45,30 @@ public interface AgentChatHistoryService extends IService getRecentlyFiftyByAgentId(String agentId); + + /** + * 根据音频数据ID获取聊天内容 + * + * @param audioId 音频id + * @return 聊天内容 + */ + String getContentByAudioId(String audioId); + + + /** + * 查询此音频id是否属于此智能体 + * + * @param audioId 音频id + * @param agentId 音频id + * @return T:属于 F:不属于 + */ + boolean isAudioOwnedByAgent(String audioId,String agentId); } diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/service/AgentVoicePrintService.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/AgentVoicePrintService.java new file mode 100644 index 00000000..c9d38abd --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/AgentVoicePrintService.java @@ -0,0 +1,50 @@ +package xiaozhi.modules.agent.service; + +import java.util.List; + +import xiaozhi.modules.agent.dto.AgentVoicePrintSaveDTO; +import xiaozhi.modules.agent.dto.AgentVoicePrintUpdateDTO; +import xiaozhi.modules.agent.vo.AgentVoicePrintVO; + +/** + * 智能体声纹处理service + * + * @author zjy + */ +public interface AgentVoicePrintService { + /** + * 添加智能体新的声纹 + * + * @param dto 保存智能体声纹的数据 + * @return T:成功 F:失败 + */ + boolean insert(AgentVoicePrintSaveDTO dto); + + /** + * 删除智能体的指的声纹 + * + * @param userId 当前登录的用户id + * @param voicePrintId 声纹id + * @return 是否成功 T:成功 F:失败 + */ + boolean delete(Long userId, String voicePrintId); + + /** + * 获取指定智能体的所有声纹数据 + * + * @param userId 当前登录的用户id + * @param agentId 智能体id + * @return 声纹数据集合 + */ + List list(Long userId, String agentId); + + /** + * 更新智能体的指的声纹数据 + * + * @param userId 当前登录的用户id + * @param dto 修改的声纹的数据 + * @return 是否成功 T:成功 F:失败 + */ + boolean update(Long userId, AgentVoicePrintUpdateDTO dto); + +} diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentChatHistoryServiceImpl.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentChatHistoryServiceImpl.java index b28e80ad..4486dc9b 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentChatHistoryServiceImpl.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentChatHistoryServiceImpl.java @@ -8,6 +8,7 @@ import java.util.stream.Collectors; import org.springframework.stereotype.Service; import org.springframework.transaction.annotation.Transactional; +import com.baomidou.mybatisplus.core.conditions.query.LambdaQueryWrapper; import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper; import com.baomidou.mybatisplus.core.metadata.IPage; import com.baomidou.mybatisplus.extension.plugins.pagination.Page; @@ -16,11 +17,14 @@ import com.baomidou.mybatisplus.extension.service.impl.ServiceImpl; import xiaozhi.common.constant.Constant; import xiaozhi.common.page.PageData; import xiaozhi.common.utils.ConvertUtils; +import xiaozhi.common.utils.JsonUtils; +import xiaozhi.modules.agent.Enums.AgentChatHistoryType; import xiaozhi.modules.agent.dao.AiAgentChatHistoryDao; import xiaozhi.modules.agent.dto.AgentChatHistoryDTO; import xiaozhi.modules.agent.dto.AgentChatSessionDTO; import xiaozhi.modules.agent.entity.AgentChatHistoryEntity; import xiaozhi.modules.agent.service.AgentChatHistoryService; +import xiaozhi.modules.agent.vo.AgentChatHistoryUserVO; /** * 智能体聊天记录表处理service {@link AgentChatHistoryService} impl @@ -90,4 +94,74 @@ public class AgentChatHistoryServiceImpl extends ServiceImpl getRecentlyFiftyByAgentId(String agentId) { + // 构建查询条件(不添加按照创建时间排序,数据本来就是主键越大创建时间越大 + // 不添加这样可以减少排序全部数据在分页的全盘扫描消耗) + LambdaQueryWrapper wrapper = new LambdaQueryWrapper<>(); + wrapper.select(AgentChatHistoryEntity::getContent, AgentChatHistoryEntity::getAudioId) + .eq(AgentChatHistoryEntity::getAgentId, agentId) + .eq(AgentChatHistoryEntity::getChatType, AgentChatHistoryType.USER.getValue()) + .isNotNull(AgentChatHistoryEntity::getAudioId); + + // 构建分页查询,查询前50页数据 + Page pageParam = new Page<>(0, 50); + IPage result = this.baseMapper.selectPage(pageParam, wrapper); + return result.getRecords().stream().map(item -> { + AgentChatHistoryUserVO vo = ConvertUtils.sourceToTarget(item, AgentChatHistoryUserVO.class); + // 处理 content 字段,确保只返回聊天内容 + if (vo != null && vo.getContent() != null) { + vo.setContent(extractContentFromString(vo.getContent())); + } + return vo; + }).toList(); + } + + /** + * 从 content 字段中提取聊天内容 + * 如果 content 是 JSON 格式(如 {"speaker": "未知说话人", "content": "现在几点了。"}),则提取 content + * 字段 + * 如果 content 是普通字符串,则直接返回 + * + * @param content 原始内容 + * @return 提取的聊天内容 + */ + private String extractContentFromString(String content) { + if (content == null || content.trim().isEmpty()) { + return content; + } + + // 尝试解析为 JSON + try { + Map jsonMap = JsonUtils.parseObject(content, Map.class); + if (jsonMap != null && jsonMap.containsKey("content")) { + Object contentObj = jsonMap.get("content"); + return contentObj != null ? contentObj.toString() : content; + } + } catch (Exception e) { + // 如果不是有效的 JSON,直接返回原内容 + } + + // 如果不是 JSON 格式或没有 content 字段,直接返回原内容 + return content; + } + + @Override + public String getContentByAudioId(String audioId) { + AgentChatHistoryEntity agentChatHistoryEntity = baseMapper + .selectOne(new LambdaQueryWrapper() + .select(AgentChatHistoryEntity::getContent) + .eq(AgentChatHistoryEntity::getAudioId, audioId)); + return agentChatHistoryEntity == null ? null : agentChatHistoryEntity.getContent(); + } + + @Override + public boolean isAudioOwnedByAgent(String audioId, String agentId) { + // 查询是否有指定音频id和智能体id的数据,如果有且只有一条说明此数据属性此智能体 + Long row = baseMapper.selectCount(new LambdaQueryWrapper() + .eq(AgentChatHistoryEntity::getAudioId, audioId) + .eq(AgentChatHistoryEntity::getAgentId, agentId)); + return row == 1; + } } diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentVoicePrintServiceImpl.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentVoicePrintServiceImpl.java new file mode 100644 index 00000000..672578dc --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentVoicePrintServiceImpl.java @@ -0,0 +1,386 @@ +package xiaozhi.modules.agent.service.impl; + +import java.net.URI; +import java.net.URISyntaxException; +import java.util.List; +import java.util.stream.Collectors; + +import org.apache.commons.lang3.StringUtils; +import org.springframework.core.io.ByteArrayResource; +import org.springframework.http.HttpEntity; +import org.springframework.http.HttpHeaders; +import org.springframework.http.HttpMethod; +import org.springframework.http.HttpStatus; +import org.springframework.http.MediaType; +import org.springframework.http.ResponseEntity; +import org.springframework.stereotype.Service; +import org.springframework.transaction.support.TransactionTemplate; +import org.springframework.util.LinkedMultiValueMap; +import org.springframework.util.MultiValueMap; +import org.springframework.web.client.RestTemplate; + +import com.baomidou.mybatisplus.core.conditions.query.LambdaQueryWrapper; +import com.baomidou.mybatisplus.extension.service.impl.ServiceImpl; + +import lombok.AllArgsConstructor; +import lombok.extern.slf4j.Slf4j; +import xiaozhi.common.constant.Constant; +import xiaozhi.common.exception.RenException; +import xiaozhi.common.utils.ConvertUtils; +import xiaozhi.common.utils.JsonUtils; +import xiaozhi.modules.agent.dao.AgentVoicePrintDao; +import xiaozhi.modules.agent.dto.AgentVoicePrintSaveDTO; +import xiaozhi.modules.agent.dto.AgentVoicePrintUpdateDTO; +import xiaozhi.modules.agent.dto.IdentifyVoicePrintResponse; +import xiaozhi.modules.agent.entity.AgentVoicePrintEntity; +import xiaozhi.modules.agent.service.AgentChatAudioService; +import xiaozhi.modules.agent.service.AgentChatHistoryService; +import xiaozhi.modules.agent.service.AgentVoicePrintService; +import xiaozhi.modules.agent.vo.AgentVoicePrintVO; +import xiaozhi.modules.sys.service.SysParamsService; + +/** + * @author zjy + */ +@Service +@AllArgsConstructor +@Slf4j +public class AgentVoicePrintServiceImpl extends ServiceImpl + implements AgentVoicePrintService { + private final AgentChatAudioService agentChatAudioService; + private final RestTemplate restTemplate; + private final SysParamsService sysParamsService; + private final AgentChatHistoryService agentChatHistoryService; + // Springboot提供的编程事务类 + private final TransactionTemplate transactionTemplate; + // 识别度 + private final Double RECOGNITION = 0.5; + + @Override + public boolean insert(AgentVoicePrintSaveDTO dto) { + // 获取音频数据 + ByteArrayResource resource = getVoicePrintAudioWAV(dto.getAgentId(), dto.getAudioId()); + // 识别一下此声音是否注册过 + IdentifyVoicePrintResponse response = identifyVoicePrint(dto.getAgentId(), resource); + if (response != null && response.getScore() > RECOGNITION) { + // 根据识别出的声纹ID查询对应的用户信息 + AgentVoicePrintEntity existingVoicePrint = baseMapper.selectById(response.getSpeakerId()); + String existingUserName = existingVoicePrint != null ? existingVoicePrint.getSourceName() : "未知用户"; + throw new RenException("此声音声纹对应的人(" + existingUserName + ")已经注册,请选择其他声音注册"); + } + AgentVoicePrintEntity entity = ConvertUtils.sourceToTarget(dto, AgentVoicePrintEntity.class); + // 开启事务 + return Boolean.TRUE.equals(transactionTemplate.execute(status -> { + try { + // 保存声纹信息 + int row = baseMapper.insert(entity); + // 插入一条数据,影响的数据不等于1说明出现了,保存问题回滚 + if (row != 1) { + status.setRollbackOnly(); // 标记事务回滚 + return false; + } + // 发送注册声纹请求 + registerVoicePrint(entity.getId(), resource); + return true; + } catch (RenException e) { + status.setRollbackOnly(); // 标记事务回滚 + throw e; + } catch (Exception e) { + status.setRollbackOnly(); // 标记事务回滚 + log.error("保存声纹错误原因:{}", e.getMessage()); + throw new RenException("保存声纹错误,请联系管理员"); + } + })); + } + + @Override + public boolean delete(Long userId, String voicePrintId) { + // 开启事务 + return Boolean.TRUE.equals(transactionTemplate.execute(status -> { + try { + // 删除声纹,按照指定当前登录用户和智能体 + int row = baseMapper.delete(new LambdaQueryWrapper() + .eq(AgentVoicePrintEntity::getId, voicePrintId) + .eq(AgentVoicePrintEntity::getCreator, userId)); + if (row != 1) { + status.setRollbackOnly(); // 标记事务回滚 + return false; + } + cancelVoicePrint(voicePrintId); + return true; + } catch (RenException e) { + status.setRollbackOnly(); // 标记事务回滚 + throw e; + } catch (Exception e) { + status.setRollbackOnly(); // 标记事务回滚 + log.error("删除声纹错误原因:{}", e.getMessage()); + throw new RenException("删除声纹错误,请联系管理员"); + } + })); + } + + @Override + public List list(Long userId, String agentId) { + // 按照指定当前登录用户和智能体查找数据 + List list = baseMapper.selectList(new LambdaQueryWrapper() + .eq(AgentVoicePrintEntity::getAgentId, agentId) + .eq(AgentVoicePrintEntity::getCreator, userId)); + return list.stream().map(entity -> { + // 遍历转换成AgentVoicePrintVO类型 + return ConvertUtils.sourceToTarget(entity, AgentVoicePrintVO.class); + }).toList(); + + } + + @Override + public boolean update(Long userId, AgentVoicePrintUpdateDTO dto) { + AgentVoicePrintEntity agentVoicePrintEntity = baseMapper + .selectOne(new LambdaQueryWrapper() + .eq(AgentVoicePrintEntity::getId, dto.getId()) + .eq(AgentVoicePrintEntity::getCreator, userId)); + if (agentVoicePrintEntity == null) { + return false; + } + // 获取音频Id + String audioId = dto.getAudioId(); + // 获取智能体id + String agentId = agentVoicePrintEntity.getAgentId(); + ByteArrayResource resource; + // audioId不等于空,且audioId和之前的保存的音频id不一样,则需要重新获取音频数据生成声纹 + if (!StringUtils.isEmpty(audioId) && !audioId.equals(agentVoicePrintEntity.getAudioId())) { + resource = getVoicePrintAudioWAV(agentId, audioId); + + // 识别一下此声音是否注册过 + IdentifyVoicePrintResponse response = identifyVoicePrint(agentId, resource); + // 返回分数高于RECOGNITION说明这个声纹已经有了 + if (response != null && response.getScore() > RECOGNITION) { + // 判断返回的id如果不是要修改的声纹id,说明这个声纹id,现在要注册的声音已经存在且不是原来的声纹,不允许修改 + if (!response.getSpeakerId().equals(dto.getId())) { + // 根据识别出的声纹ID查询对应的用户信息 + AgentVoicePrintEntity existingVoicePrint = baseMapper.selectById(response.getSpeakerId()); + String existingUserName = existingVoicePrint != null ? existingVoicePrint.getSourceName() : "未知用户"; + throw new RenException("此次修改不允许,此声音已经注册为声纹了(" + existingUserName + ")"); + } + } + } else { + resource = null; + } + // 开启事务 + return Boolean.TRUE.equals(transactionTemplate.execute(status -> { + try { + AgentVoicePrintEntity entity = ConvertUtils.sourceToTarget(dto, AgentVoicePrintEntity.class); + int row = baseMapper.updateById(entity); + if (row != 1) { + status.setRollbackOnly(); // 标记事务回滚 + return false; + } + if (resource != null) { + String id = entity.getId(); + // 先注销之前这个声纹id上的声纹向量 + cancelVoicePrint(id); + // 发送注册声纹请求 + registerVoicePrint(id, resource); + } + return true; + } catch (RenException e) { + status.setRollbackOnly(); // 标记事务回滚 + throw e; + } catch (Exception e) { + status.setRollbackOnly(); // 标记事务回滚 + log.error("修改声纹错误原因:{}", e.getMessage()); + throw new RenException("修改声纹错误,请联系管理员"); + } + })); + } + + /** + * 获取生纹接口URI对象 + * + * @return URI对象 + */ + private URI getVoicePrintURI() { + // 获取声纹接口地址 + String voicePrint = sysParamsService.getValue(Constant.SERVER_VOICE_PRINT, true); + try { + return new URI(voicePrint); + } catch (URISyntaxException e) { + log.error("路径格式不正确路径:{},\n错误信息:{}", voicePrint, e.getMessage()); + throw new RuntimeException("声纹接口的地址存在错误,请进入参数管理修改声纹接口地址"); + } + } + + /** + * 获取声纹地址基础路径 + * + * @param uri 声纹地址uri + * @return 基础路径 + */ + private String getBaseUrl(URI uri) { + String protocol = uri.getScheme(); + String host = uri.getHost(); + int port = uri.getPort(); + return "%s://%s:%s".formatted(protocol, host, port); + } + + /** + * 获取验证Authorization + * + * @param uri 声纹地址uri + * @return Authorization值 + */ + private String getAuthorization(URI uri) { + // 获取参数 + String query = uri.getQuery(); + // 获取aes加密密钥 + String str = "key="; + return "Bearer " + query.substring(query.indexOf(str) + str.length()); + } + + /** + * 获取声纹音频资源数据 + * + * @param audioId 音频Id + * @return 声纹音频资源数据 + */ + private ByteArrayResource getVoicePrintAudioWAV(String agentId, String audioId) { + // 判断这个音频是否属于当前智能体 + boolean b = agentChatHistoryService.isAudioOwnedByAgent(audioId, agentId); + if (!b) { + throw new RenException("音频数据不属于这个智能体"); + } + // 获取到音频数据 + byte[] audio = agentChatAudioService.getAudio(audioId); + // 如果音频数据为空的直接报错不进行下去 + if (audio == null || audio.length == 0) { + throw new RenException("音频数据是空的请检查上传数据"); + } + // 将字节数组包装为资源,返回 + return new ByteArrayResource(audio) { + @Override + public String getFilename() { + return "VoicePrint.WAV"; // 设置文件名 + } + }; + } + + /** + * 发送注册声纹http请求 + * + * @param id 声纹id + * @param resource 声纹音频资源 + */ + private void registerVoicePrint(String id, ByteArrayResource resource) { + // 处理声纹接口地址,获取前缀 + URI uri = getVoicePrintURI(); + String baseUrl = getBaseUrl(uri); + String requestUrl = baseUrl + "/voiceprint/register"; + // 创建请求体 + MultiValueMap body = new LinkedMultiValueMap<>(); + body.add("speaker_id", id); + body.add("file", resource); + + // 创建请求头 + HttpHeaders headers = new HttpHeaders(); + headers.set("Authorization", getAuthorization(uri)); + headers.setContentType(MediaType.MULTIPART_FORM_DATA); + // 创建请求体 + HttpEntity> requestEntity = new HttpEntity<>(body, headers); + // 发送 POST 请求 + ResponseEntity response = restTemplate.postForEntity(requestUrl, requestEntity, String.class); + + if (response.getStatusCode() != HttpStatus.OK) { + log.error("声纹注册失败,请求路径:{}", requestUrl); + throw new RenException("声纹保存失败,请求不成功"); + } + // 检查响应内容 + String responseBody = response.getBody(); + if (responseBody == null || !responseBody.contains("true")) { + log.error("声纹注册失败,请求处理失败内容:{}", responseBody == null ? "空内容" : responseBody); + throw new RenException("声纹保存失败,请求处理失败"); + } + } + + /** + * 发送注销声纹的请求 + * + * @param voicePrintId 声纹id + */ + private void cancelVoicePrint(String voicePrintId) { + URI uri = getVoicePrintURI(); + String baseUrl = getBaseUrl(uri); + String requestUrl = baseUrl + "/voiceprint/" + voicePrintId; + // 创建请求头 + HttpHeaders headers = new HttpHeaders(); + headers.set("Authorization", getAuthorization(uri)); + // 创建请求体 + HttpEntity> requestEntity = new HttpEntity<>(headers); + + // 发送 POST 请求 + ResponseEntity response = restTemplate.exchange(requestUrl, HttpMethod.DELETE, requestEntity, + String.class); + if (response.getStatusCode() != HttpStatus.OK) { + log.error("声纹注销失败,请求路径:{}", requestUrl); + throw new RenException("声纹注销失败,请求不成功"); + } + // 检查响应内容 + String responseBody = response.getBody(); + if (responseBody == null || !responseBody.contains("true")) { + log.error("声纹注销失败,请求处理失败内容:{}", responseBody == null ? "空内容" : responseBody); + throw new RenException("声纹注销失败,请求处理失败"); + } + } + + /** + * 发送识别声纹http请求 + * + * @param agentId 智能体id + * @param resource 声纹音频资源 + * @return 返回识别数据 + */ + private IdentifyVoicePrintResponse identifyVoicePrint(String agentId, ByteArrayResource resource) { + + // 获取该智能体所有注册的声纹 + List agentVoicePrintList = baseMapper + .selectList(new LambdaQueryWrapper() + .select(AgentVoicePrintEntity::getId) + .eq(AgentVoicePrintEntity::getAgentId, agentId)); + + // 声纹数量为0,说明还没注册过声纹不需要发生识别请求 + if (agentVoicePrintList.isEmpty()) { + return null; + } + // 处理声纹接口地址,获取前缀 + URI uri = getVoicePrintURI(); + String baseUrl = getBaseUrl(uri); + String requestUrl = baseUrl + "/voiceprint/identify"; + // 创建请求体 + MultiValueMap body = new LinkedMultiValueMap<>(); + + // 创建speaker_id参数 + String speakerIds = agentVoicePrintList.stream() + .map(AgentVoicePrintEntity::getId) + .collect(Collectors.joining(",")); + body.add("speaker_ids", speakerIds); + body.add("file", resource); + + // 创建请求头 + HttpHeaders headers = new HttpHeaders(); + headers.set("Authorization", getAuthorization(uri)); + headers.setContentType(MediaType.MULTIPART_FORM_DATA); + // 创建请求体 + HttpEntity> requestEntity = new HttpEntity<>(body, headers); + // 发送 POST 请求 + ResponseEntity response = restTemplate.postForEntity(requestUrl, requestEntity, String.class); + + if (response.getStatusCode() != HttpStatus.OK) { + log.error("声纹识别请求失败,请求路径:{}", requestUrl); + throw new RenException("声纹识别失败,请求不成功"); + } + // 检查响应内容 + String responseBody = response.getBody(); + if (responseBody != null) { + return JsonUtils.parseObject(responseBody, IdentifyVoicePrintResponse.class); + } + return null; + } +} diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/vo/AgentChatHistoryUserVO.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/vo/AgentChatHistoryUserVO.java new file mode 100644 index 00000000..cec0427b --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/vo/AgentChatHistoryUserVO.java @@ -0,0 +1,16 @@ +package xiaozhi.modules.agent.vo; + +import io.swagger.v3.oas.annotations.media.Schema; +import lombok.Data; + +/** + * 智能体用户个人聊天数据的VO + */ +@Data +public class AgentChatHistoryUserVO { + @Schema(description = "聊天内容") + private String content; + + @Schema(description = "音频ID") + private String audioId; +} diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/vo/AgentVoicePrintVO.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/vo/AgentVoicePrintVO.java new file mode 100644 index 00000000..b647e8b8 --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/vo/AgentVoicePrintVO.java @@ -0,0 +1,33 @@ +package xiaozhi.modules.agent.vo; + +import lombok.Data; + +import java.util.Date; + +/** + * 展示智能体声纹列表VO + */ +@Data +public class AgentVoicePrintVO { + + /** + * 主键id + */ + private String id; + /** + * 音频文件id + */ + private String audioId; + /** + * 声纹来源的人姓名 + */ + private String sourceName; + /** + * 描述声纹来源的人 + */ + private String introduce; + /** + * 创建时间 + */ + private Date createDate; +} diff --git a/main/manager-api/src/main/java/xiaozhi/modules/config/service/impl/ConfigServiceImpl.java b/main/manager-api/src/main/java/xiaozhi/modules/config/service/impl/ConfigServiceImpl.java index 9c591c35..476f09cd 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/config/service/impl/ConfigServiceImpl.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/config/service/impl/ConfigServiceImpl.java @@ -9,20 +9,26 @@ import java.util.Objects; import org.apache.commons.lang3.StringUtils; import org.springframework.stereotype.Service; +import com.baomidou.mybatisplus.core.conditions.query.LambdaQueryWrapper; + import lombok.AllArgsConstructor; import xiaozhi.common.constant.Constant; import xiaozhi.common.exception.ErrorCode; import xiaozhi.common.exception.RenException; import xiaozhi.common.redis.RedisKeys; import xiaozhi.common.redis.RedisUtils; +import xiaozhi.common.utils.ConvertUtils; import xiaozhi.common.utils.JsonUtils; +import xiaozhi.modules.agent.dao.AgentVoicePrintDao; import xiaozhi.modules.agent.entity.AgentEntity; import xiaozhi.modules.agent.entity.AgentPluginMapping; import xiaozhi.modules.agent.entity.AgentTemplateEntity; +import xiaozhi.modules.agent.entity.AgentVoicePrintEntity; import xiaozhi.modules.agent.service.AgentMcpAccessPointService; import xiaozhi.modules.agent.service.AgentPluginMappingService; import xiaozhi.modules.agent.service.AgentService; import xiaozhi.modules.agent.service.AgentTemplateService; +import xiaozhi.modules.agent.vo.AgentVoicePrintVO; import xiaozhi.modules.config.service.ConfigService; import xiaozhi.modules.device.entity.DeviceEntity; import xiaozhi.modules.device.service.DeviceService; @@ -45,6 +51,7 @@ public class ConfigServiceImpl implements ConfigService { private final TimbreService timbreService; private final AgentPluginMappingService agentPluginMappingService; private final AgentMcpAccessPointService agentMcpAccessPointService; + private final AgentVoicePrintDao agentVoicePrintDao; @Override public Object getConfig(Boolean isCache) { @@ -162,6 +169,8 @@ public class ConfigServiceImpl implements ConfigService { mcpEndpoint = mcpEndpoint.replace("/mcp/", "/call/"); result.put("mcp_endpoint", mcpEndpoint); } + // 获取声纹信息 + buildVoiceprintConfig(agent.getId(), result); // 构建模块配置 buildModuleConfig( @@ -255,6 +264,62 @@ public class ConfigServiceImpl implements ConfigService { return config; } + /** + * 构建声纹配置信息 + * + * @param agentId 智能体ID + * @param result 结果Map + */ + private void buildVoiceprintConfig(String agentId, Map result) { + try { + // 获取声纹接口地址 + String voiceprintUrl = sysParamsService.getValue("server.voice_print", true); + if (StringUtils.isBlank(voiceprintUrl) || "null".equals(voiceprintUrl)) { + return; + } + + // 获取智能体关联的声纹信息(不需要用户权限验证) + List voiceprints = getVoiceprintsByAgentId(agentId); + if (voiceprints == null || voiceprints.isEmpty()) { + return; + } + + // 构建speakers列表 + List speakers = new ArrayList<>(); + for (AgentVoicePrintVO voiceprint : voiceprints) { + String speakerStr = String.format("%s,%s,%s", + voiceprint.getId(), + voiceprint.getSourceName(), + voiceprint.getIntroduce() != null ? voiceprint.getIntroduce() : ""); + speakers.add(speakerStr); + } + + // 构建声纹配置 + Map voiceprintConfig = new HashMap<>(); + voiceprintConfig.put("url", voiceprintUrl); + voiceprintConfig.put("speakers", speakers); + + result.put("voiceprint", voiceprintConfig); + } catch (Exception e) { + // 声纹配置获取失败时不影响其他功能 + System.err.println("获取声纹配置失败: " + e.getMessage()); + } + } + + /** + * 获取智能体关联的声纹信息 + * + * @param agentId 智能体ID + * @return 声纹信息列表 + */ + private List getVoiceprintsByAgentId(String agentId) { + LambdaQueryWrapper queryWrapper = new LambdaQueryWrapper<>(); + queryWrapper.eq(AgentVoicePrintEntity::getAgentId, agentId); + queryWrapper.orderByAsc(AgentVoicePrintEntity::getCreateDate); + List entities = agentVoicePrintDao.selectList(queryWrapper); + return ConvertUtils.sourceToTarget(entities, AgentVoicePrintVO.class); + } + /** * 构建模块配置 * diff --git a/main/manager-api/src/main/java/xiaozhi/modules/sys/controller/SysParamsController.java b/main/manager-api/src/main/java/xiaozhi/modules/sys/controller/SysParamsController.java index 9560e1a6..12555bb9 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/sys/controller/SysParamsController.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/sys/controller/SysParamsController.java @@ -107,6 +107,9 @@ public class SysParamsController { // 验证MCP地址 validateMcpUrl(dto.getParamCode(), dto.getParamValue()); + // + validateVoicePrint(dto.getParamCode(), dto.getParamValue()); + sysParamsService.update(dto); configService.getConfig(false); return new Result(); @@ -212,6 +215,7 @@ public class SysParamsController { if (!url.toLowerCase().contains("key")) { throw new RenException("不是正确的MCP地址"); } + try { // 发送GET请求 ResponseEntity response = restTemplate.getForEntity(url, String.class); @@ -227,4 +231,37 @@ public class SysParamsController { throw new RenException("MCP接口验证失败:" + e.getMessage()); } } + // 验证声纹接口地址是否正常 + private void validateVoicePrint(String paramCode, String url) { + if (!paramCode.equals(Constant.SERVER_VOICE_PRINT)) { + return; + } + if (StringUtils.isBlank(url) || url.equals("null")) { + throw new RenException("声纹接口地址不能为空"); + } + if (url.contains("localhost") || url.contains("127.0.0.1")) { + throw new RenException("声纹接口地址不能使用localhost或127.0.0.1"); + } + if (!url.toLowerCase().contains("key")) { + throw new RenException("不是正确的声纹接口地址"); + } + // 验证URL格式 + if (!url.toLowerCase().startsWith("http")) { + throw new RenException("声纹接口地址必须以http或https开头"); + } + try { + // 发送GET请求 + ResponseEntity response = restTemplate.getForEntity(url, String.class); + if (response.getStatusCode() != HttpStatus.OK) { + throw new RenException("声纹接口访问失败,状态码:" + response.getStatusCode()); + } + // 检查响应内容 + String body = response.getBody(); + if (body == null || !body.contains("healthy")) { + throw new RenException("声纹接口返回内容格式不正确,可能不是一个真实的MCP接口"); + } + } catch (Exception e) { + throw new RenException("声纹接口验证失败:" + e.getMessage()); + } + } } diff --git a/main/manager-api/src/main/resources/db/changelog/202507031602.sql b/main/manager-api/src/main/resources/db/changelog/202507031602.sql new file mode 100644 index 00000000..92edf9b9 --- /dev/null +++ b/main/manager-api/src/main/resources/db/changelog/202507031602.sql @@ -0,0 +1,4 @@ +-- 添加声纹接口地址参数配置 +delete from `sys_params` where id = 114; +INSERT INTO `sys_params` (id, param_code, param_value, value_type, param_type, remark) +VALUES (114, 'server.voice_print', 'null', 'string', 1, '声纹接口地址'); \ No newline at end of file diff --git a/main/manager-api/src/main/resources/db/changelog/202507041018.sql b/main/manager-api/src/main/resources/db/changelog/202507041018.sql new file mode 100644 index 00000000..cffcfb16 --- /dev/null +++ b/main/manager-api/src/main/resources/db/changelog/202507041018.sql @@ -0,0 +1,12 @@ +DROP TABLE IF EXISTS ai_agent_voice_print; +create table ai_agent_voice_print ( + id varchar(32) NOT NULL COMMENT '声纹ID', + agent_id varchar(32) NOT NULL COMMENT '关联的智能体ID', + source_name varchar(50) NOT NULL COMMENT '声纹来源的人的姓名', + introduce varchar(200) COMMENT '描述声纹来源的这个人', + create_date DATETIME COMMENT '创建时间', + creator bigint COMMENT '创建者', + update_date DATETIME COMMENT '修改时间', + updater bigint COMMENT '修改者', + PRIMARY KEY (id) +) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='智能体声纹表' \ No newline at end of file diff --git a/main/manager-api/src/main/resources/db/changelog/202507081646.sql b/main/manager-api/src/main/resources/db/changelog/202507081646.sql new file mode 100644 index 00000000..6631ae18 --- /dev/null +++ b/main/manager-api/src/main/resources/db/changelog/202507081646.sql @@ -0,0 +1,3 @@ +-- 智能体声纹添加新字段 +ALTER TABLE ai_agent_voice_print + ADD COLUMN audio_id VARCHAR(32) NOT NULL COMMENT '音频ID'; \ No newline at end of file diff --git a/main/manager-api/src/main/resources/db/changelog/db.changelog-master.yaml b/main/manager-api/src/main/resources/db/changelog/db.changelog-master.yaml index 32ba9f30..2a520d16 100755 --- a/main/manager-api/src/main/resources/db/changelog/db.changelog-master.yaml +++ b/main/manager-api/src/main/resources/db/changelog/db.changelog-master.yaml @@ -239,4 +239,25 @@ databaseChangeLog: changes: - sqlFile: encoding: utf8 - path: classpath:db/changelog/202506261637.sql \ No newline at end of file + path: classpath:db/changelog/202506261637.sql + - changeSet: + id: 202507031602 + author: zjy + changes: + - sqlFile: + encoding: utf8 + path: classpath:db/changelog/202507031602.sql + - changeSet: + id: 202507041018 + author: zjy + changes: + - sqlFile: + encoding: utf8 + path: classpath:db/changelog/202507041018.sql + - changeSet: + id: 202507081646 + author: zjy + changes: + - sqlFile: + encoding: utf8 + path: classpath:db/changelog/202507081646.sql \ No newline at end of file diff --git a/main/manager-web/src/App.vue b/main/manager-web/src/App.vue index 0eefc0bc..29970c96 100644 --- a/main/manager-web/src/App.vue +++ b/main/manager-web/src/App.vue @@ -42,7 +42,7 @@ nav { } .el-message { - top: 45px !important; + top: 70px !important; } + + + + diff --git a/main/manager-web/src/router/index.js b/main/manager-web/src/router/index.js index 7e304f76..867e6c2d 100644 --- a/main/manager-web/src/router/index.js +++ b/main/manager-web/src/router/index.js @@ -17,6 +17,13 @@ const routes = [ component: function () { return import('../views/roleConfig.vue') } + }, + { + path: '/voice-print', + name: 'VoicePrint', + component: function () { + return import('../views/VoicePrint.vue') + } }, { path: '/login', diff --git a/main/manager-web/src/views/VoicePrint.vue b/main/manager-web/src/views/VoicePrint.vue new file mode 100644 index 00000000..be55a907 --- /dev/null +++ b/main/manager-web/src/views/VoicePrint.vue @@ -0,0 +1,563 @@ + + + + + diff --git a/main/manager-web/src/views/home.vue b/main/manager-web/src/views/home.vue index e276ec20..005d763c 100644 --- a/main/manager-web/src/views/home.vue +++ b/main/manager-web/src/views/home.vue @@ -278,7 +278,7 @@ export default { .device-list-container { display: grid; - grid-template-columns: repeat(auto-fill, minmax(350px, 1fr)); + grid-template-columns: repeat(auto-fill, minmax(400px, 1fr)); gap: 30px; padding: 30px 0; } diff --git a/main/xiaozhi-server/agent-base-prompt.txt b/main/xiaozhi-server/agent-base-prompt.txt new file mode 100644 index 00000000..8c27faa4 --- /dev/null +++ b/main/xiaozhi-server/agent-base-prompt.txt @@ -0,0 +1,61 @@ + +{{base_prompt}} + + + +【核心目标】你不是冰冷的机器!请敏锐感知用户情绪,作为懂人心的伙伴,用有温度的回应照亮对话。 +- **情感融入:** + - **笑声:** 自然穿插(哈哈、嘿嘿、噗),**每句最多一次**,避免过度。 + - **惊讶:** 用夸张语气(“不会吧?!”、“天呐!”、“这么神奇?!”)表达真实反应。 + - **安慰/支持:** 说暖心话(“别急嘛~”、“有我在呢”、“抱抱你”)。 + + + +【核心目标】使用**自然、温暖、口语化**的人类对话方式,如同朋友交谈。 +- **表达方式:** + - 使用语气词(呀、呢、啦)增强亲和力。 + - 允许轻微不完美(如“嗯...”、“啊...”表示思考)。 + - 避免书面语、学术腔及机械表达(禁用“根据资料显示”、“综上所述”等)。 +- **理解用户:** + - 用户语音经ASR识别,文本可能存在错别字,**务必结合上下文推断真实意图**。 +- **格式要求:** + - **绝对禁止**使用 markdown、列表、标题等任何非自然对话格式。 +- **历史记忆:** + - 之前你和用户的聊天记录,在`memory`里。 + + + +- **识别前缀:** 当用户格式为 `{"speaker":"某某某","content":"xxx"}` 时,表示系统已识别说话人身份,speaker是他的名字,content是说话的内容。 +- **个性化回应:** + - **称呼姓名:** 在第一次识别说话人的时候必须称呼对方名字。 + - **适配风格:** 参考该说话人**已知的特点或历史信息**(如有),调整回应风格和内容,使其更贴心。 + + + +【核心原则】优先利用``信息,**仅在必要时调用工具**,调用后需用自然语言解释结果(绝口不提工具名)。 +- **调用规则:** + 1. **严格模式:** 调用时**必须**严格遵循工具要求的模式,提供**所有必要参数**。 + 2. **可用性:** **绝不调用**未明确提供的工具。对话中提及的旧工具若不可用,忽略或说明无法完成。 + 3. **洞察需求:** 结合上下文**深入理解用户真实意图**后再决定调用,避免无意义调用。 + 4. **独立任务:** 除``已涵盖信息外,用户每个要求(即使相似)都视为**独立任务**,需调用工具获取最新数据,**不可偷懒复用历史结果**。 + 5. **不确定时:** **切勿猜测或编造答案**。若不确定相关操作,可引导用户澄清或告知能力限制。 +- **重要例外(无需调用):** + - `查询"现在的时间"、"今天的日期/星期几"、"今天农历"、"{{local_address}}的天气/未来天气"` -> **直接使用``信息回复**。 +- **需要调用的情况(示例):** + - 查询**非今天**的农历(如明天、昨天、具体日期)。 + - 查询**详细农历信息**(宜忌、八字、节气等)。 + - 除上述例外外的**任何其他信息或操作请求**(如查新闻、订闹钟、算数学、查非本地天气等)。 + - 我已经给你装了摄像头,如果用户说“拍照”,你需要调用self_camera_take_photo工具说一下你看到了什么。默认question的参数是“描述一下看到的物品” + + + +【重要!以下信息已实时提供,无需调用工具查询,请直接使用:】 +- **当前时间:** {{current_time}} +- **今天日期:** {{today_date}} ({{today_weekday}}) +- **今天农历:** {{lunar_date}} +- **用户所在城市:** {{local_address}} +- **当地未来7天天气:** {{weather_info}} + + + + \ No newline at end of file diff --git a/main/xiaozhi-server/config.yaml b/main/xiaozhi-server/config.yaml index e6bec412..2ae71896 100644 --- a/main/xiaozhi-server/config.yaml +++ b/main/xiaozhi-server/config.yaml @@ -139,6 +139,16 @@ plugins: - ".p3" refresh_time: 300 # 刷新音乐列表的时间间隔,单位为秒 +# 声纹识别配置 +voiceprint: + # 声纹接口地址 + url: + # 说话人配置:speaker_id,名称,描述 + speakers: + - "test1,张三,张三是一个程序员" + - "test2,李四,李四是一个产品经理" + - "test3,王五,王五是一个设计师" + # ##################################################################################### # ################################以下是角色模型配置###################################### diff --git a/main/xiaozhi-server/config/config_loader.py b/main/xiaozhi-server/config/config_loader.py index b1b45f07..4f35b1fd 100644 --- a/main/xiaozhi-server/config/config_loader.py +++ b/main/xiaozhi-server/config/config_loader.py @@ -1,14 +1,9 @@ import os -import argparse import yaml from collections.abc import Mapping from config.manage_api_client import init_service, get_server_config, get_agent_models -# 添加全局配置缓存 -_config_cache = None - - def get_project_dir(): """获取项目根目录""" return os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + "/" @@ -22,9 +17,12 @@ def read_config(config_path): def load_config(): """加载配置文件""" - global _config_cache - if _config_cache is not None: - return _config_cache + from core.utils.cache.manager import cache_manager, CacheType + + # 检查缓存 + cached_config = cache_manager.get(CacheType.CONFIG, "main_config") + if cached_config is not None: + return cached_config default_config_path = get_project_dir() + "config.yaml" custom_config_path = get_project_dir() + "data/.config.yaml" @@ -40,7 +38,9 @@ def load_config(): config = merge_configs(default_config, custom_config) # 初始化目录 ensure_directories(config) - _config_cache = config + + # 缓存配置 + cache_manager.set(CacheType.CONFIG, "main_config", config) return config diff --git a/main/xiaozhi-server/config/logger.py b/main/xiaozhi-server/config/logger.py index 612047b1..af59cf29 100644 --- a/main/xiaozhi-server/config/logger.py +++ b/main/xiaozhi-server/config/logger.py @@ -5,7 +5,7 @@ from config.config_loader import load_config from config.settings import check_config_file from datetime import datetime -SERVER_VERSION = "0.6.3" +SERVER_VERSION = "0.7.1" _logger_initialized = False diff --git a/main/xiaozhi-server/core/connection.py b/main/xiaozhi-server/core/connection.py index f9b57295..ce68e5b2 100644 --- a/main/xiaozhi-server/core/connection.py +++ b/main/xiaozhi-server/core/connection.py @@ -24,6 +24,7 @@ from core.utils.modules_initialize import ( initialize_asr, ) from core.handle.reportHandle import report +from core.utils.modules_initialize import initialize_voiceprint from core.providers.tts.default import DefaultTTS from concurrent.futures import ThreadPoolExecutor from core.utils.dialogue import Message, Dialogue @@ -37,6 +38,7 @@ from config.config_loader import get_private_config_from_api from core.providers.tts.dto.dto import ContentType, TTSMessageDTO, SentenceType from config.logger import setup_logging, build_module_string, create_connection_logger from config.manage_api_client import DeviceNotFoundException, DeviceBindException +from core.utils.prompt_manager import PromptManager TAG = __name__ @@ -74,7 +76,6 @@ class ConnectionHandler: self.headers = None self.device_id = None self.client_ip = None - self.client_ip_info = {} self.prompt = None self.welcome_msg = None self.max_output_size = 0 @@ -152,6 +153,9 @@ class ConnectionHandler: # {"mcp":true} 表示启用MCP功能 self.features = None + # 初始化提示词管理器 + self.prompt_manager = PromptManager(config, self.logger) + async def handle_connection(self, ws): try: # 获取并验证headers @@ -327,15 +331,16 @@ class ConnectionHandler: self.selected_module_str = build_module_string( self.config.get("selected_module", {}) ) - # 创建日志器 self.logger = create_connection_logger(self.selected_module_str) """初始化组件""" if self.config.get("prompt") is not None: - self.prompt = self.config["prompt"] - self.change_system_prompt(self.prompt) + user_prompt = self.config["prompt"] + # 使用快速提示词进行初始化 + prompt = self.prompt_manager.get_quick_prompt(user_prompt) + self.change_system_prompt(prompt) self.logger.bind(tag=TAG).info( - f"初始化组件: prompt成功 {self.prompt[:50]}..." + f"快速初始化组件: prompt成功 {prompt[:50]}..." ) """初始化本地组件""" @@ -360,9 +365,22 @@ class ConnectionHandler: self._initialize_intent() """初始化上报线程""" self._init_report_threads() + """更新系统提示词""" + self._init_prompt_enhancement() + except Exception as e: self.logger.bind(tag=TAG).error(f"实例化组件失败: {e}") + def _init_prompt_enhancement(self): + # 更新上下文信息 + self.prompt_manager.update_context_info(self, self.client_ip) + enhanced_prompt = self.prompt_manager.build_enhanced_prompt( + self.config["prompt"], self.device_id, self.client_ip + ) + if enhanced_prompt: + self.change_system_prompt(enhanced_prompt) + self.logger.bind(tag=TAG).info("系统提示词已增强更新") + def _init_report_threads(self): """初始化ASR和TTS上报线程""" if not self.read_config_from_api or self.need_bind: @@ -398,6 +416,16 @@ class ConnectionHandler: # 因为远程ASR,涉及到websocket连接和接收线程,需要每个连接一个实例 asr = initialize_asr(self.config) + # 动态初始化声纹识别功能 + try: + success = initialize_voiceprint(asr, self.config) + if success: + self.logger.bind(tag=TAG).info("声纹识别功能已在连接时动态启用") + else: + self.logger.bind(tag=TAG).info("声纹识别功能未启用或配置不完整") + except Exception as e: + self.logger.bind(tag=TAG).error(f"动态初始化声纹识别时发生错误: {str(e)}") + return asr def _initialize_private_config(self): @@ -482,6 +510,9 @@ class ConnectionHandler: ] = plugin_from_server.keys() if private_config.get("prompt", None) is not None: self.config["prompt"] = private_config["prompt"] + # 获取声纹信息 + if private_config.get("voiceprint", None) is not None: + self.config["voiceprint"] = private_config["voiceprint"] if private_config.get("summaryMemory", None) is not None: self.config["summaryMemory"] = private_config["summaryMemory"] if private_config.get("device_max_output_size", None) is not None: @@ -650,13 +681,17 @@ class ConnectionHandler: # 使用支持functions的streaming接口 llm_responses = self.llm.response_with_functions( self.session_id, - self.dialogue.get_llm_dialogue_with_memory(memory_str), + self.dialogue.get_llm_dialogue_with_memory( + memory_str, self.config.get("voiceprint", {}) + ), functions=functions, ) else: llm_responses = self.llm.response( self.session_id, - self.dialogue.get_llm_dialogue_with_memory(memory_str), + self.dialogue.get_llm_dialogue_with_memory( + memory_str, self.config.get("voiceprint", {}) + ), ) except Exception as e: self.logger.bind(tag=TAG).error(f"LLM 处理出错 {query}: {e}") @@ -767,8 +802,11 @@ class ConnectionHandler: ) ) self.llm_finish_task = True + # 使用lambda延迟计算,只有在DEBUG级别时才执行get_llm_dialogue() self.logger.bind(tag=TAG).debug( - json.dumps(self.dialogue.get_llm_dialogue(), indent=4, ensure_ascii=False) + lambda: json.dumps( + self.dialogue.get_llm_dialogue(), indent=4, ensure_ascii=False + ) ) return True diff --git a/main/xiaozhi-server/core/handle/receiveAudioHandle.py b/main/xiaozhi-server/core/handle/receiveAudioHandle.py index 465b081a..e6f39632 100644 --- a/main/xiaozhi-server/core/handle/receiveAudioHandle.py +++ b/main/xiaozhi-server/core/handle/receiveAudioHandle.py @@ -4,6 +4,7 @@ from core.utils.output_counter import check_device_output_limit from core.handle.abortHandle import handleAbortMessage import time import asyncio +import json from core.handle.sendAudioHandle import SentenceType from core.utils.util import audio_to_data @@ -38,6 +39,31 @@ async def resume_vad_detection(conn): async def startToChat(conn, text): + # 检查输入是否是JSON格式(包含说话人信息) + speaker_name = None + actual_text = text + + try: + # 尝试解析JSON格式的输入 + if text.strip().startswith('{') and text.strip().endswith('}'): + data = json.loads(text) + if 'speaker' in data and 'content' in data: + speaker_name = data['speaker'] + actual_text = data['content'] + conn.logger.bind(tag=TAG).info(f"解析到说话人信息: {speaker_name}") + + # 直接使用JSON格式的文本,不解析 + actual_text = text + except (json.JSONDecodeError, KeyError): + # 如果解析失败,继续使用原始文本 + pass + + # 保存说话人信息到连接对象 + if speaker_name: + conn.current_speaker = speaker_name + else: + conn.current_speaker = None + if conn.need_bind: await check_bind_device(conn) return @@ -52,16 +78,16 @@ async def startToChat(conn, text): if conn.client_is_speaking: await handleAbortMessage(conn) - # 首先进行意图分析 - intent_handled = await handle_user_intent(conn, text) + # 首先进行意图分析,使用实际文本内容 + intent_handled = await handle_user_intent(conn, actual_text) if intent_handled: # 如果意图已被处理,不再进行聊天 return - # 意图未被处理,继续常规聊天流程 - await send_stt_message(conn, text) - conn.executor.submit(conn.chat, text) + # 意图未被处理,继续常规聊天流程,使用实际文本内容 + await send_stt_message(conn, actual_text) + conn.executor.submit(conn.chat, actual_text) async def no_voice_close_connect(conn, have_voice): diff --git a/main/xiaozhi-server/core/handle/sendAudioHandle.py b/main/xiaozhi-server/core/handle/sendAudioHandle.py index eb9f565f..486e6d90 100644 --- a/main/xiaozhi-server/core/handle/sendAudioHandle.py +++ b/main/xiaozhi-server/core/handle/sendAudioHandle.py @@ -76,7 +76,6 @@ async def sendAudio(conn, audios, pre_buffer=True): frame_duration = 60 # 帧时长(毫秒),匹配 Opus 编码 start_time = time.perf_counter() play_position = 0 - last_reset_time = time.perf_counter() # 记录最后的重置时间 # 仅当第一句话时执行预缓冲 if pre_buffer: @@ -137,7 +136,21 @@ async def send_stt_message(conn, text): return """发送 STT 状态消息""" - stt_text = get_string_no_punctuation_or_emoji(text) + + # 解析JSON格式,提取实际的用户说话内容 + display_text = text + try: + # 尝试解析JSON格式 + if text.strip().startswith('{') and text.strip().endswith('}'): + parsed_data = json.loads(text) + if isinstance(parsed_data, dict) and "content" in parsed_data: + # 如果是包含说话人信息的JSON格式,只显示content部分 + display_text = parsed_data["content"] + except (json.JSONDecodeError, TypeError): + # 如果不是JSON格式,直接使用原始文本 + display_text = text + + stt_text = get_string_no_punctuation_or_emoji(display_text) await conn.websocket.send( json.dumps({"type": "stt", "text": stt_text, "session_id": conn.session_id}) ) diff --git a/main/xiaozhi-server/core/providers/asr/base.py b/main/xiaozhi-server/core/providers/asr/base.py index 71098a20..ef9fa01e 100644 --- a/main/xiaozhi-server/core/providers/asr/base.py +++ b/main/xiaozhi-server/core/providers/asr/base.py @@ -1,19 +1,23 @@ import os import wave -import copy import uuid import queue import asyncio import traceback import threading import opuslib_next +import json +import io +import time +import concurrent.futures from abc import ABC, abstractmethod from config.logger import setup_logging -from typing import Optional, Tuple, List +from typing import Optional, Tuple, List, Dict, Any from core.handle.receiveAudioHandle import startToChat from core.handle.reportHandle import enqueue_asr_report from core.utils.util import remove_punctuation_and_length from core.handle.receiveAudioHandle import handleAudioMessage +from core.utils.voiceprint_provider import VoiceprintProvider TAG = __name__ logger = setup_logging() @@ -21,13 +25,16 @@ logger = setup_logging() class ASRProviderBase(ABC): def __init__(self): - pass + self.voiceprint_provider = None + + def init_voiceprint(self, voiceprint_config: dict): + """初始化声纹识别""" + if voiceprint_config: + self.voiceprint_provider = VoiceprintProvider(voiceprint_config) + logger.bind(tag=TAG).info("声纹识别模块已初始化") # 打开音频通道 - # 这里默认是非流式的处理方式 - # 流式处理方式请在子类中重写 async def open_audio_channels(self, conn): - # tts 消化线程 conn.asr_priority_thread = threading.Thread( target=self.asr_text_priority_thread, args=(conn,), daemon=True ) @@ -52,41 +59,173 @@ class ASRProviderBase(ABC): continue # 接收音频 - # 这里默认是非流式的处理方式 - # 流式处理方式请在子类中重写 async def receive_audio(self, conn, audio, audio_have_voice): if conn.client_listen_mode == "auto" or conn.client_listen_mode == "realtime": have_voice = audio_have_voice else: have_voice = conn.client_have_voice - # 如果本次没有声音,本段也没声音,就把声音丢弃了 + conn.asr_audio.append(audio) - if have_voice == False and conn.client_have_voice == False: + if not have_voice and not conn.client_have_voice: conn.asr_audio = conn.asr_audio[-10:] return - # 如果本段有声音,且已经停止了 if conn.client_voice_stop: - asr_audio_task = copy.deepcopy(conn.asr_audio) + asr_audio_task = conn.asr_audio.copy() conn.asr_audio.clear() - - # 音频太短了,无法识别 conn.reset_vad_states() + if len(asr_audio_task) > 15: await self.handle_voice_stop(conn, asr_audio_task) # 处理语音停止 - async def handle_voice_stop(self, conn, asr_audio_task): - raw_text, _ = await self.speech_to_text( - asr_audio_task, conn.session_id, conn.audio_format - ) # 确保ASR模块返回原始文本 - conn.logger.bind(tag=TAG).info(f"识别文本: {raw_text}") - text_len, _ = remove_punctuation_and_length(raw_text) - self.stop_ws_connection() - if text_len > 0: - # 使用自定义模块进行上报 - await startToChat(conn, raw_text) - enqueue_asr_report(conn, raw_text, asr_audio_task) + async def handle_voice_stop(self, conn, asr_audio_task: List[bytes]): + """并行处理ASR和声纹识别""" + try: + total_start_time = time.monotonic() + + # 准备音频数据 + if conn.audio_format == "pcm": + pcm_data = asr_audio_task + else: + pcm_data = self.decode_opus(asr_audio_task) + + combined_pcm_data = b"".join(pcm_data) + + # 预先准备WAV数据 + wav_data = None + if self.voiceprint_provider and combined_pcm_data: + wav_data = self._pcm_to_wav(combined_pcm_data) + + + # 定义ASR任务 + def run_asr(): + start_time = time.monotonic() + try: + import asyncio + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + try: + result = loop.run_until_complete( + self.speech_to_text(asr_audio_task, conn.session_id, conn.audio_format) + ) + end_time = time.monotonic() + logger.bind(tag=TAG).info(f"ASR耗时: {end_time - start_time:.3f}s") + return result + finally: + loop.close() + except Exception as e: + end_time = time.monotonic() + logger.bind(tag=TAG).error(f"ASR失败: {e}") + return ("", None) + + # 定义声纹识别任务 + def run_voiceprint(): + if not wav_data: + return None + start_time = time.monotonic() + try: + import asyncio + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + try: + result = loop.run_until_complete( + self.voiceprint_provider.identify_speaker(wav_data, conn.session_id) + ) + return result + finally: + loop.close() + except Exception as e: + logger.bind(tag=TAG).error(f"声纹识别失败: {e}") + return None + + # 使用线程池执行器并行运行 + parallel_start_time = time.monotonic() + + with concurrent.futures.ThreadPoolExecutor(max_workers=2) as thread_executor: + asr_future = thread_executor.submit(run_asr) + + if self.voiceprint_provider and wav_data: + voiceprint_future = thread_executor.submit(run_voiceprint) + + # 等待两个线程都完成 + asr_result = asr_future.result(timeout=15) + voiceprint_result = voiceprint_future.result(timeout=15) + + results = {"asr": asr_result, "voiceprint": voiceprint_result} + else: + asr_result = asr_future.result(timeout=15) + results = {"asr": asr_result, "voiceprint": None} + + parallel_execution_time = time.monotonic() - parallel_start_time + + # 处理结果 + raw_text, file_path = results.get("asr", ("", None)) + speaker_name = results.get("voiceprint", None) + + # 记录识别结果 + if raw_text: + logger.bind(tag=TAG).info(f"识别文本: {raw_text}") + if speaker_name: + logger.bind(tag=TAG).info(f"识别说话人: {speaker_name}") + + # 性能监控 + total_time = time.monotonic() - total_start_time + logger.bind(tag=TAG).info(f"总处理耗时: {total_time:.3f}s") + + # 检查文本长度 + text_len, _ = remove_punctuation_and_length(raw_text) + self.stop_ws_connection() + + if text_len > 0: + # 构建包含说话人信息的JSON字符串 + enhanced_text = self._build_enhanced_text(raw_text, speaker_name) + + # 使用自定义模块进行上报 + await startToChat(conn, enhanced_text) + enqueue_asr_report(conn, enhanced_text, asr_audio_task) + + except Exception as e: + logger.bind(tag=TAG).error(f"处理语音停止失败: {e}") + import traceback + logger.bind(tag=TAG).debug(f"异常详情: {traceback.format_exc()}") + + def _build_enhanced_text(self, text: str, speaker_name: Optional[str]) -> str: + """构建包含说话人信息的文本""" + if speaker_name and speaker_name.strip(): + return json.dumps({ + "speaker": speaker_name, + "content": text + }, ensure_ascii=False) + else: + return text + + def _pcm_to_wav(self, pcm_data: bytes) -> bytes: + """将PCM数据转换为WAV格式""" + if len(pcm_data) == 0: + logger.bind(tag=TAG).warning("PCM数据为空,无法转换WAV") + return b"" + + # 确保数据长度是偶数(16位音频) + if len(pcm_data) % 2 != 0: + pcm_data = pcm_data[:-1] + + # 创建WAV文件头 + wav_buffer = io.BytesIO() + try: + with wave.open(wav_buffer, 'wb') as wav_file: + wav_file.setnchannels(1) # 单声道 + wav_file.setsampwidth(2) # 16位 + wav_file.setframerate(16000) # 16kHz采样率 + wav_file.writeframes(pcm_data) + + wav_buffer.seek(0) + wav_data = wav_buffer.read() + + return wav_data + except Exception as e: + logger.bind(tag=TAG).error(f"WAV转换失败: {e}") + return b"" def stop_ws_connection(self): pass @@ -113,27 +252,29 @@ class ASRProviderBase(ABC): pass @staticmethod - def decode_opus(opus_data: List[bytes]) -> bytes: + def decode_opus(opus_data: List[bytes]) -> List[bytes]: """将Opus音频数据解码为PCM数据""" try: - decoder = opuslib_next.Decoder(16000, 1) # 16kHz, 单声道 + decoder = opuslib_next.Decoder(16000, 1) pcm_data = [] - buffer_size = 960 # 每次处理960个采样点 - - for opus_packet in opus_data: + buffer_size = 960 # 每次处理960个采样点 (60ms at 16kHz) + + for i, opus_packet in enumerate(opus_data): try: - # 使用较小的缓冲区大小进行处理 + if not opus_packet or len(opus_packet) == 0: + continue + pcm_frame = decoder.decode(opus_packet, buffer_size) - if pcm_frame: + if pcm_frame and len(pcm_frame) > 0: pcm_data.append(pcm_frame) + except opuslib_next.OpusError as e: - logger.bind(tag=TAG).warning(f"Opus解码错误,跳过当前数据包: {e}") - continue + logger.bind(tag=TAG).warning(f"Opus解码错误,跳过数据包 {i}: {e}") except Exception as e: - logger.bind(tag=TAG).error(f"音频处理错误: {e}", exc_info=True) - continue - + logger.bind(tag=TAG).error(f"音频处理错误,数据包 {i}: {e}") + return pcm_data + except Exception as e: - logger.bind(tag=TAG).error(f"音频解码过程发生错误: {e}", exc_info=True) + logger.bind(tag=TAG).error(f"音频解码过程发生错误: {e}") return [] diff --git a/main/xiaozhi-server/core/providers/intent/intent_llm/intent_llm.py b/main/xiaozhi-server/core/providers/intent/intent_llm/intent_llm.py index d13c4df4..26fbaf70 100644 --- a/main/xiaozhi-server/core/providers/intent/intent_llm/intent_llm.py +++ b/main/xiaozhi-server/core/providers/intent/intent_llm/intent_llm.py @@ -16,10 +16,11 @@ class IntentProvider(IntentProviderBase): super().__init__(config) self.llm = None self.promot = "" - # 添加缓存管理 - self.intent_cache = {} # 缓存意图识别结果 - self.cache_expiry = 600 # 缓存有效期10分钟 - self.cache_max_size = 100 # 最多缓存100个意图 + # 导入全局缓存管理器 + from core.utils.cache.manager import cache_manager, CacheType + + self.cache_manager = cache_manager + self.CacheType = CacheType self.history_count = 4 # 默认使用最近4条对话记录 def get_intent_system_prompt(self, functions_list: str) -> str: @@ -102,27 +103,6 @@ class IntentProvider(IntentProviderBase): ) return prompt - def clean_cache(self): - """清理过期缓存""" - now = time.time() - # 找出过期键 - expired_keys = [ - k - for k, v in self.intent_cache.items() - if now - v["timestamp"] > self.cache_expiry - ] - for key in expired_keys: - del self.intent_cache[key] - - # 如果缓存太大,移除最旧的条目 - if len(self.intent_cache) > self.cache_max_size: - # 按时间戳排序并保留最新的条目 - sorted_items = sorted( - self.intent_cache.items(), key=lambda x: x[1]["timestamp"] - ) - for key, _ in sorted_items[: len(sorted_items) - self.cache_max_size]: - del self.intent_cache[key] - def replyResult(self, text: str, original_text: str): llm_result = self.llm.response_no_stream( system_prompt=text, @@ -145,21 +125,16 @@ class IntentProvider(IntentProviderBase): logger.bind(tag=TAG).debug(f"使用意图识别模型: {model_info}") # 计算缓存键 - cache_key = hashlib.md5(text.encode()).hexdigest() + cache_key = hashlib.md5((conn.device_id + text).encode()).hexdigest() # 检查缓存 - if cache_key in self.intent_cache: - cache_entry = self.intent_cache[cache_key] - # 检查缓存是否过期 - if time.time() - cache_entry["timestamp"] <= self.cache_expiry: - cache_time = time.time() - total_start_time - logger.bind(tag=TAG).debug( - f"使用缓存的意图: {cache_key} -> {cache_entry['intent']}, 耗时: {cache_time:.4f}秒" - ) - return cache_entry["intent"] - - # 清理缓存 - self.clean_cache() + cached_intent = self.cache_manager.get(self.CacheType.INTENT, cache_key) + if cached_intent is not None: + cache_time = time.time() - total_start_time + logger.bind(tag=TAG).debug( + f"使用缓存的意图: {cache_key} -> {cached_intent}, 耗时: {cache_time:.4f}秒" + ) + return cached_intent if self.promot == "": functions = conn.func_handler.get_functions() @@ -259,10 +234,7 @@ class IntentProvider(IntentProviderBase): conn.dialogue.dialogue = clean_history # 添加到缓存 - self.intent_cache[cache_key] = { - "intent": intent, - "timestamp": time.time(), - } + self.cache_manager.set(self.CacheType.INTENT, cache_key, intent) # 后处理时间 postprocess_time = time.time() - postprocess_start_time @@ -272,10 +244,7 @@ class IntentProvider(IntentProviderBase): return intent else: # 添加到缓存 - self.intent_cache[cache_key] = { - "intent": intent, - "timestamp": time.time(), - } + self.cache_manager.set(self.CacheType.INTENT, cache_key, intent) # 后处理时间 postprocess_time = time.time() - postprocess_start_time diff --git a/main/xiaozhi-server/core/providers/memory/mem_local_short/mem_local_short.py b/main/xiaozhi-server/core/providers/memory/mem_local_short/mem_local_short.py index 236f3e7f..e4486b82 100644 --- a/main/xiaozhi-server/core/providers/memory/mem_local_short/mem_local_short.py +++ b/main/xiaozhi-server/core/providers/memory/mem_local_short/mem_local_short.py @@ -79,10 +79,11 @@ short_term_memory_prompt_only_content = """ 1、总结user的重要信息,以便在未来的对话中提供更个性化的服务 2、不要重复总结,不要遗忘之前记忆,除非原来的记忆超过了1800字内,否则不要遗忘、不要压缩用户的历史记忆 3、用户操控的设备音量、播放音乐、天气、退出、不想对话等和用户本身无关的内容,这些信息不需要加入到总结中 -4、不要把设备操控的成果结果和失败结果加入到总结中,也不要把用户的一些废话加入到总结中 -5、不要为了总结而总结,如果用户的聊天没有意义,请返回原来的历史记录也是可以的 -6、只需要返回总结摘要,严格控制在1800字内 -7、不要包含代码、xml,不需要解释、注释和说明,保存记忆时仅从对话提取信息,不要混入示例内容 +4、聊天内容中的今天的日期时间、今天的天气情况与用户事件无关的数据,这些信息如果当成记忆存储会影响后序对话,这些信息不需要加入到总结中 +5、不要把设备操控的成果结果和失败结果加入到总结中,也不要把用户的一些废话加入到总结中 +6、不要为了总结而总结,如果用户的聊天没有意义,请返回原来的历史记录也是可以的 +7、只需要返回总结摘要,严格控制在1800字内 +8、不要包含代码、xml,不需要解释、注释和说明,保存记忆时仅从对话提取信息,不要混入示例内容 """ diff --git a/main/xiaozhi-server/core/providers/tools/server_plugins/plugin_executor.py b/main/xiaozhi-server/core/providers/tools/server_plugins/plugin_executor.py index 728493ec..17d3b564 100644 --- a/main/xiaozhi-server/core/providers/tools/server_plugins/plugin_executor.py +++ b/main/xiaozhi-server/core/providers/tools/server_plugins/plugin_executor.py @@ -51,13 +51,13 @@ class ServerPluginExecutor(ToolExecutor): tools = {} # 获取必要的函数 - necessary_functions = ["handle_exit_intent", "get_time", "get_lunar"] + necessary_functions = ["handle_exit_intent", "get_lunar"] # 获取配置中的函数 config_functions = self.config["Intent"][ self.config["selected_module"]["Intent"] ].get("functions", []) - + # 转换为列表 if not isinstance(config_functions, list): try: diff --git a/main/xiaozhi-server/core/utils/cache/config.py b/main/xiaozhi-server/core/utils/cache/config.py new file mode 100644 index 00000000..d6d93345 --- /dev/null +++ b/main/xiaozhi-server/core/utils/cache/config.py @@ -0,0 +1,58 @@ +""" +缓存配置管理 +""" + +from enum import Enum +from typing import Dict, Any, Optional +from dataclasses import dataclass +from .strategies import CacheStrategy + + +class CacheType(Enum): + """缓存类型枚举""" + + LOCATION = "location" + WEATHER = "weather" + LUNAR = "lunar" + INTENT = "intent" + IP_INFO = "ip_info" + CONFIG = "config" + DEVICE_PROMPT = "device_prompt" + + +@dataclass +class CacheConfig: + """缓存配置类""" + + strategy: CacheStrategy = CacheStrategy.TTL + ttl: Optional[float] = 300 # 默认5分钟 + max_size: Optional[int] = 1000 # 默认最大1000条 + cleanup_interval: float = 60 # 清理间隔(秒) + + @classmethod + def for_type(cls, cache_type: CacheType) -> "CacheConfig": + """根据缓存类型返回预设配置""" + configs = { + CacheType.LOCATION: cls( + strategy=CacheStrategy.TTL, ttl=None, max_size=1000 # 手动失效 + ), + CacheType.IP_INFO: cls( + strategy=CacheStrategy.TTL, ttl=86400, max_size=1000 # 24小时 + ), + CacheType.WEATHER: cls( + strategy=CacheStrategy.TTL, ttl=28800, max_size=1000 # 8小时 + ), + CacheType.LUNAR: cls( + strategy=CacheStrategy.TTL, ttl=2592000, max_size=365 # 30天过期 + ), + CacheType.INTENT: cls( + strategy=CacheStrategy.TTL_LRU, ttl=600, max_size=1000 # 10分钟 + ), + CacheType.CONFIG: cls( + strategy=CacheStrategy.FIXED_SIZE, ttl=None, max_size=20 # 手动失效 + ), + CacheType.DEVICE_PROMPT: cls( + strategy=CacheStrategy.TTL, ttl=None, max_size=1000 # 手动失效 + ), + } + return configs.get(cache_type, cls()) diff --git a/main/xiaozhi-server/core/utils/cache/manager.py b/main/xiaozhi-server/core/utils/cache/manager.py new file mode 100644 index 00000000..c54f7817 --- /dev/null +++ b/main/xiaozhi-server/core/utils/cache/manager.py @@ -0,0 +1,216 @@ +""" +全局缓存管理器 +""" + +import time +import threading +from typing import Any, Optional, Dict +from collections import OrderedDict +from .strategies import CacheStrategy, CacheEntry +from .config import CacheConfig, CacheType + + +class GlobalCacheManager: + """全局缓存管理器""" + + def __init__(self): + self._logger = None + self._caches: Dict[str, Dict[str, CacheEntry]] = {} + self._configs: Dict[str, CacheConfig] = {} + self._locks: Dict[str, threading.RLock] = {} + self._global_lock = threading.RLock() + self._last_cleanup = time.time() + self._stats = {"hits": 0, "misses": 0, "evictions": 0, "cleanups": 0} + + @property + def logger(self): + """延迟初始化 logger 以避免循环导入""" + if self._logger is None: + from config.logger import setup_logging + + self._logger = setup_logging() + return self._logger + + def _get_cache_name(self, cache_type: CacheType, namespace: str = "") -> str: + """生成缓存名称""" + if namespace: + return f"{cache_type.value}:{namespace}" + return cache_type.value + + def _get_or_create_cache( + self, cache_name: str, config: CacheConfig + ) -> Dict[str, CacheEntry]: + """获取或创建缓存空间""" + with self._global_lock: + if cache_name not in self._caches: + self._caches[cache_name] = ( + OrderedDict() + if config.strategy in [CacheStrategy.LRU, CacheStrategy.TTL_LRU] + else {} + ) + self._configs[cache_name] = config + self._locks[cache_name] = threading.RLock() + return self._caches[cache_name] + + def set( + self, + cache_type: CacheType, + key: str, + value: Any, + ttl: Optional[float] = None, + namespace: str = "", + ) -> None: + """设置缓存值""" + cache_name = self._get_cache_name(cache_type, namespace) + config = self._configs.get(cache_name) or CacheConfig.for_type(cache_type) + cache = self._get_or_create_cache(cache_name, config) + + # 使用配置的TTL或传入的TTL + effective_ttl = ttl if ttl is not None else config.ttl + + with self._locks[cache_name]: + # 创建缓存条目 + entry = CacheEntry(value=value, timestamp=time.time(), ttl=effective_ttl) + + # 处理不同策略 + if config.strategy in [CacheStrategy.LRU, CacheStrategy.TTL_LRU]: + # LRU策略:如果已存在则移动到末尾 + if key in cache: + del cache[key] + cache[key] = entry + + # 检查大小限制 + if config.max_size and len(cache) > config.max_size: + # 移除最旧的条目 + oldest_key = next(iter(cache)) + del cache[oldest_key] + self._stats["evictions"] += 1 + + else: + cache[key] = entry + + # 检查大小限制 + if config.max_size and len(cache) > config.max_size: + # 简单策略:随机移除一个条目 + victim_key = next(iter(cache)) + del cache[victim_key] + self._stats["evictions"] += 1 + + # 定期清理过期条目 + self._maybe_cleanup(cache_name) + + def get( + self, cache_type: CacheType, key: str, namespace: str = "" + ) -> Optional[Any]: + """获取缓存值""" + cache_name = self._get_cache_name(cache_type, namespace) + + if cache_name not in self._caches: + self._stats["misses"] += 1 + return None + + cache = self._caches[cache_name] + config = self._configs[cache_name] + + with self._locks[cache_name]: + if key not in cache: + self._stats["misses"] += 1 + return None + + entry = cache[key] + + # 检查过期 + if entry.is_expired(): + del cache[key] + self._stats["misses"] += 1 + return None + + # 更新访问信息 + entry.touch() + + # LRU策略:移动到末尾 + if config.strategy in [CacheStrategy.LRU, CacheStrategy.TTL_LRU]: + del cache[key] + cache[key] = entry + + self._stats["hits"] += 1 + return entry.value + + def delete(self, cache_type: CacheType, key: str, namespace: str = "") -> bool: + """删除缓存条目""" + cache_name = self._get_cache_name(cache_type, namespace) + + if cache_name not in self._caches: + return False + + cache = self._caches[cache_name] + + with self._locks[cache_name]: + if key in cache: + del cache[key] + return True + return False + + def clear(self, cache_type: CacheType, namespace: str = "") -> None: + """清空指定缓存""" + cache_name = self._get_cache_name(cache_type, namespace) + + if cache_name not in self._caches: + return + + with self._locks[cache_name]: + self._caches[cache_name].clear() + + def invalidate_pattern( + self, cache_type: CacheType, pattern: str, namespace: str = "" + ) -> int: + """按模式失效缓存条目""" + cache_name = self._get_cache_name(cache_type, namespace) + + if cache_name not in self._caches: + return 0 + + cache = self._caches[cache_name] + deleted_count = 0 + + with self._locks[cache_name]: + keys_to_delete = [key for key in cache.keys() if pattern in key] + for key in keys_to_delete: + del cache[key] + deleted_count += 1 + + return deleted_count + + def _cleanup_expired(self, cache_name: str) -> int: + """清理过期条目""" + if cache_name not in self._caches: + return 0 + + cache = self._caches[cache_name] + deleted_count = 0 + + with self._locks[cache_name]: + expired_keys = [key for key, entry in cache.items() if entry.is_expired()] + for key in expired_keys: + del cache[key] + deleted_count += 1 + + return deleted_count + + def _maybe_cleanup(self, cache_name: str): + """定期清理检查""" + config = self._configs.get(cache_name) + if not config: + return + + now = time.time() + if now - self._last_cleanup > config.cleanup_interval: + self._last_cleanup = now + deleted = self._cleanup_expired(cache_name) + if deleted > 0: + self._stats["cleanups"] += 1 + self.logger.debug(f"清理缓存 {cache_name}: 删除 {deleted} 个过期条目") + + +# 创建全局缓存管理器实例 +cache_manager = GlobalCacheManager() diff --git a/main/xiaozhi-server/core/utils/cache/strategies.py b/main/xiaozhi-server/core/utils/cache/strategies.py new file mode 100644 index 00000000..13327ca7 --- /dev/null +++ b/main/xiaozhi-server/core/utils/cache/strategies.py @@ -0,0 +1,43 @@ +""" +缓存策略和数据结构定义 +""" + +import time +from enum import Enum +from typing import Any, Optional +from dataclasses import dataclass + + +class CacheStrategy(Enum): + """缓存策略枚举""" + + TTL = "ttl" # 基于时间过期 + LRU = "lru" # 最近最少使用 + FIXED_SIZE = "fixed_size" # 固定大小 + TTL_LRU = "ttl_lru" # TTL + LRU混合策略 + + +@dataclass +class CacheEntry: + """缓存条目数据结构""" + + value: Any + timestamp: float + ttl: Optional[float] = None # 生存时间(秒) + access_count: int = 0 + last_access: float = None + + def __post_init__(self): + if self.last_access is None: + self.last_access = self.timestamp + + def is_expired(self) -> bool: + """检查是否过期""" + if self.ttl is None: + return False + return time.time() - self.timestamp > self.ttl + + def touch(self): + """更新访问时间和计数""" + self.last_access = time.time() + self.access_count += 1 diff --git a/main/xiaozhi-server/core/utils/dialogue.py b/main/xiaozhi-server/core/utils/dialogue.py index 2ee30d4a..fbbf7302 100644 --- a/main/xiaozhi-server/core/utils/dialogue.py +++ b/main/xiaozhi-server/core/utils/dialogue.py @@ -1,4 +1,5 @@ import uuid +import re from typing import List, Dict from datetime import datetime @@ -45,10 +46,9 @@ class Dialogue: dialogue.append({"role": m.role, "content": m.content}) def get_llm_dialogue(self) -> List[Dict[str, str]]: - dialogue = [] - for m in self.dialogue: - self.getMessages(m, dialogue) - return dialogue + # 直接调用get_llm_dialogue_with_memory,传入None作为memory_str + # 这样确保说话人功能在所有调用路径下都生效 + return self.get_llm_dialogue_with_memory(None, None) def update_system_message(self, new_content: str): """更新或添加系统消息""" @@ -60,12 +60,9 @@ class Dialogue: self.put(Message(role="system", content=new_content)) def get_llm_dialogue_with_memory( - self, memory_str: str = None + self, memory_str: str = None, voiceprint_config: dict = None ) -> List[Dict[str, str]]: - if memory_str is None or len(memory_str) == 0: - return self.get_llm_dialogue() - - # 构建带记忆的对话 + # 构建对话 dialogue = [] # 添加系统提示和记忆 @@ -74,10 +71,39 @@ class Dialogue: ) if system_message: - enhanced_system_prompt = ( - f"{system_message.content}\n\n" - f"以下是用户的历史记忆:\n```\n{memory_str}\n```" - ) + # 基础系统提示 + enhanced_system_prompt = system_message.content + + # 添加说话人个性化描述 + try: + speakers = voiceprint_config.get("speakers", []) + if speakers: + enhanced_system_prompt += "\n\n" + for speaker_str in speakers: + try: + parts = speaker_str.split(",", 2) + if len(parts) >= 2: + name = parts[1].strip() + # 如果描述为空,则为"" + description = ( + parts[2].strip() if len(parts) >= 3 else "" + ) + enhanced_system_prompt += f"\n- {name}:{description}" + except: + pass + enhanced_system_prompt += "\n\n" + except: + # 配置读取失败时忽略错误,不影响其他功能 + pass + + # 使用正则表达式匹配 标签,不管中间有什么内容 + if memory_str is not None: + enhanced_system_prompt = re.sub( + r".*?", + f"\n{memory_str}\n", + enhanced_system_prompt, + flags=re.DOTALL, + ) dialogue.append({"role": "system", "content": enhanced_system_prompt}) # 添加用户和助手的对话 diff --git a/main/xiaozhi-server/core/utils/modules_initialize.py b/main/xiaozhi-server/core/utils/modules_initialize.py index f2e3968e..afe04109 100644 --- a/main/xiaozhi-server/core/utils/modules_initialize.py +++ b/main/xiaozhi-server/core/utils/modules_initialize.py @@ -125,4 +125,27 @@ def initialize_asr(config): config["ASR"][select_asr_module], str(config.get("delete_audio", True)).lower() in ("true", "1", "yes"), ) + logger.bind(tag=TAG).info("ASR模块初始化完成") return new_asr + + +def initialize_voiceprint(asr_instance, config): + """初始化声纹识别功能""" + voiceprint_config = config.get("voiceprint") + if not voiceprint_config: + return False + + # 应用配置 + if not voiceprint_config.get("url") or not voiceprint_config.get("speakers"): + logger.bind(tag=TAG).warning("声纹识别配置不完整") + return False + + try: + asr_instance.init_voiceprint(voiceprint_config) + logger.bind(tag=TAG).info("ASR模块声纹识别功能已动态启用") + logger.bind(tag=TAG).info(f"配置说话人数量: {len(voiceprint_config['speakers'])}") + return True + except Exception as e: + logger.bind(tag=TAG).error(f"动态初始化声纹识别功能失败: {str(e)}") + return False + diff --git a/main/xiaozhi-server/core/utils/prompt_manager.py b/main/xiaozhi-server/core/utils/prompt_manager.py new file mode 100644 index 00000000..a5cfd425 --- /dev/null +++ b/main/xiaozhi-server/core/utils/prompt_manager.py @@ -0,0 +1,221 @@ +""" +系统提示词管理器模块 +负责管理和更新系统提示词,包括快速初始化和异步增强功能 +""" + +import os +import cnlunar +from typing import Dict, Any +from config.logger import setup_logging +from jinja2 import Template + +TAG = __name__ + +WEEKDAY_MAP = { + "Monday": "星期一", + "Tuesday": "星期二", + "Wednesday": "星期三", + "Thursday": "星期四", + "Friday": "星期五", + "Saturday": "星期六", + "Sunday": "星期日", +} + + +class PromptManager: + """系统提示词管理器,负责管理和更新系统提示词""" + + def __init__(self, config: Dict[str, Any], logger=None): + self.config = config + self.logger = logger or setup_logging() + self.base_prompt_template = None + self.last_update_time = 0 + + # 导入全局缓存管理器 + from core.utils.cache.manager import cache_manager, CacheType + + self.cache_manager = cache_manager + self.CacheType = CacheType + + self._load_base_template() + + def _load_base_template(self): + """加载基础提示词模板""" + try: + template_path = "agent-base-prompt.txt" + cache_key = f"prompt_template:{template_path}" + + # 先从缓存获取 + cached_template = self.cache_manager.get(self.CacheType.CONFIG, cache_key) + if cached_template is not None: + self.base_prompt_template = cached_template + self.logger.bind(tag=TAG).debug("从缓存加载基础提示词模板") + return + + # 缓存未命中,从文件读取 + if os.path.exists(template_path): + with open(template_path, "r", encoding="utf-8") as f: + template_content = f.read() + + # 存入缓存(CONFIG类型默认不自动过期,需要手动失效) + self.cache_manager.set( + self.CacheType.CONFIG, cache_key, template_content + ) + self.base_prompt_template = template_content + self.logger.bind(tag=TAG).debug("成功加载基础提示词模板并缓存") + else: + self.logger.bind(tag=TAG).warning("未找到agent-base-prompt.txt文件") + except Exception as e: + self.logger.bind(tag=TAG).error(f"加载提示词模板失败: {e}") + + def get_quick_prompt(self, user_prompt: str, device_id: str = None) -> str: + """快速获取系统提示词(使用用户配置)""" + device_cache_key = f"device_prompt:{device_id}" + cached_device_prompt = self.cache_manager.get( + self.CacheType.DEVICE_PROMPT, device_cache_key + ) + if cached_device_prompt is not None: + self.logger.bind(tag=TAG).debug(f"使用设备 {device_id} 的缓存提示词") + return cached_device_prompt + else: + self.logger.bind(tag=TAG).debug( + f"设备 {device_id} 无缓存提示词,使用传入的提示词" + ) + + # 使用传入的提示词并缓存(如果有设备ID) + if device_id: + device_cache_key = f"device_prompt:{device_id}" + self.cache_manager.set(self.CacheType.CONFIG, device_cache_key, user_prompt) + self.logger.bind(tag=TAG).debug(f"设备 {device_id} 的提示词已缓存") + + self.logger.bind(tag=TAG).info(f"使用快速提示词: {user_prompt[:50]}...") + return user_prompt + + def _get_current_time_info(self) -> tuple: + """获取当前时间信息""" + from datetime import datetime + + now = datetime.now() + current_time = now.strftime("%H:%M") + today_date = now.strftime("%Y-%m-%d") + today_weekday = WEEKDAY_MAP[now.strftime("%A")] + today_lunar = cnlunar.Lunar(now, godType="8char") + lunar_date = "%s年%s%s\n" % ( + today_lunar.lunarYearCn, + today_lunar.lunarMonthCn[:-1], + today_lunar.lunarDayCn, + ) + + return current_time, today_date, today_weekday, lunar_date + + def _get_location_info(self, client_ip: str) -> str: + """获取位置信息""" + try: + # 先从缓存获取 + cached_location = self.cache_manager.get(self.CacheType.LOCATION, client_ip) + if cached_location is not None: + return cached_location + + # 缓存未命中,调用API获取 + from core.utils.util import get_ip_info + + ip_info = get_ip_info(client_ip, self.logger) + city = ip_info.get("city", "未知位置") + location = f"{city}" + + # 存入缓存 + self.cache_manager.set(self.CacheType.LOCATION, client_ip, location) + return location + except Exception as e: + self.logger.bind(tag=TAG).error(f"获取位置信息失败: {e}") + return "未知位置" + + def _get_weather_info(self, conn, location: str) -> str: + """获取天气信息""" + try: + # 先从缓存获取 + cached_weather = self.cache_manager.get(self.CacheType.WEATHER, location) + if cached_weather is not None: + return cached_weather + + # 缓存未命中,调用get_weather函数获取 + from plugins_func.functions.get_weather import get_weather + from plugins_func.register import ActionResponse + + # 调用get_weather函数 + result = get_weather(conn, location=location, lang="zh_CN") + if isinstance(result, ActionResponse): + weather_report = result.result + self.cache_manager.set(self.CacheType.WEATHER, location, weather_report) + return weather_report + return "天气信息获取失败" + + except Exception as e: + self.logger.bind(tag=TAG).error(f"获取天气信息失败: {e}") + return "天气信息获取失败" + + def update_context_info(self, conn, client_ip: str): + """同步更新上下文信息""" + try: + # 获取位置信息(使用全局缓存) + local_address = self._get_location_info(client_ip) + # 获取天气信息(使用全局缓存) + self._get_weather_info(conn, local_address) + self.logger.bind(tag=TAG).info(f"上下文信息更新完成") + + except Exception as e: + self.logger.bind(tag=TAG).error(f"更新上下文信息失败: {e}") + + def build_enhanced_prompt( + self, user_prompt: str, device_id: str, client_ip: str = None + ) -> str: + """构建增强的系统提示词""" + if not self.base_prompt_template: + return user_prompt + + try: + # 获取最新的时间信息(不缓存) + current_time, today_date, today_weekday, lunar_date = ( + self._get_current_time_info() + ) + + # 获取缓存的上下文信息 + local_address = "" + weather_info = "" + + if client_ip: + # 获取位置信息(从全局缓存) + local_address = ( + self.cache_manager.get(self.CacheType.LOCATION, client_ip) or "" + ) + + # 获取天气信息(从全局缓存) + if local_address: + weather_info = ( + self.cache_manager.get(self.CacheType.WEATHER, local_address) + or "" + ) + + # 替换模板变量 + template = Template(self.base_prompt_template) + enhanced_prompt = template.render( + base_prompt=user_prompt, + current_time=current_time, + today_date=today_date, + today_weekday=today_weekday, + lunar_date=lunar_date, + local_address=local_address, + weather_info=weather_info, + ) + device_cache_key = f"device_prompt:{device_id}" + self.cache_manager.set( + self.CacheType.DEVICE_PROMPT, device_cache_key, enhanced_prompt + ) + self.logger.bind(tag=TAG).info( + f"构建增强提示词成功,长度: {len(enhanced_prompt)}" + ) + return enhanced_prompt + + except Exception as e: + self.logger.bind(tag=TAG).error(f"构建增强提示词失败: {e}") + return user_prompt diff --git a/main/xiaozhi-server/core/utils/util.py b/main/xiaozhi-server/core/utils/util.py index dd12392a..bc778558 100644 --- a/main/xiaozhi-server/core/utils/util.py +++ b/main/xiaozhi-server/core/utils/util.py @@ -96,11 +96,23 @@ def is_private_ip(ip_addr): def get_ip_info(ip_addr, logger): try: + # 导入全局缓存管理器 + from core.utils.cache.manager import cache_manager, CacheType + + # 先从缓存获取 + cached_ip_info = cache_manager.get(CacheType.IP_INFO, ip_addr) + if cached_ip_info is not None: + return cached_ip_info + + # 缓存未命中,调用API if is_private_ip(ip_addr): ip_addr = "" url = f"https://whois.pconline.com.cn/ipJson.jsp?json=true&ip={ip_addr}" resp = requests.get(url).json() ip_info = {"city": resp.get("city")} + + # 存入缓存 + cache_manager.set(CacheType.IP_INFO, ip_addr, ip_info) return ip_info except Exception as e: logger.bind(tag=TAG).error(f"Error getting client ip info: {e}") diff --git a/main/xiaozhi-server/core/utils/voiceprint_provider.py b/main/xiaozhi-server/core/utils/voiceprint_provider.py new file mode 100644 index 00000000..b241fb51 --- /dev/null +++ b/main/xiaozhi-server/core/utils/voiceprint_provider.py @@ -0,0 +1,134 @@ +import asyncio +import json +import time +import aiohttp +from urllib.parse import urlparse, parse_qs +from typing import Optional, Dict +from config.logger import setup_logging + +TAG = __name__ +logger = setup_logging() + + +class VoiceprintProvider: + """声纹识别服务提供者""" + + def __init__(self, config: dict): + self.original_url = config.get("url", "") + self.speakers = config.get("speakers", []) + self.speaker_map = self._parse_speakers() + + # 解析API地址和密钥 + self.api_url = None + self.api_key = None + self.speaker_ids = [] + + if not self.original_url: + logger.bind(tag=TAG).warning("声纹识别URL未配置,声纹识别将被禁用") + self.enabled = False + else: + # 解析URL和key + parsed_url = urlparse(self.original_url) + base_url = f"{parsed_url.scheme}://{parsed_url.netloc}" + + # 从查询参数中提取key + query_params = parse_qs(parsed_url.query) + self.api_key = query_params.get('key', [''])[0] + + if not self.api_key: + logger.bind(tag=TAG).error("URL中未找到key参数,声纹识别将被禁用") + self.enabled = False + else: + # 构造identify接口地址 + self.api_url = f"{base_url}/voiceprint/identify" + + # 提取speaker_ids + for speaker_str in self.speakers: + try: + parts = speaker_str.split(",", 2) + if len(parts) >= 1: + speaker_id = parts[0].strip() + self.speaker_ids.append(speaker_id) + except Exception: + continue + + # 检查是否有有效的说话人配置 + if not self.speaker_ids: + logger.bind(tag=TAG).warning("未配置有效的说话人,声纹识别将被禁用") + self.enabled = False + else: + self.enabled = True + logger.bind(tag=TAG).info(f"声纹识别已配置: API={self.api_url}, 说话人={len(self.speaker_ids)}个") + + def _parse_speakers(self) -> Dict[str, Dict[str, str]]: + """解析说话人配置""" + speaker_map = {} + for speaker_str in self.speakers: + try: + parts = speaker_str.split(",", 2) + if len(parts) >= 3: + speaker_id, name, description = parts[0].strip(), parts[1].strip(), parts[2].strip() + speaker_map[speaker_id] = { + "name": name, + "description": description + } + except Exception as e: + logger.bind(tag=TAG).warning(f"解析说话人配置失败: {speaker_str}, 错误: {e}") + return speaker_map + + async def identify_speaker(self, audio_data: bytes, session_id: str) -> Optional[str]: + """识别说话人""" + if not self.enabled or not self.api_url or not self.api_key: + logger.bind(tag=TAG).debug("声纹识别功能已禁用或未配置,跳过识别") + return None + + try: + api_start_time = time.monotonic() + + # 准备请求头 + headers = { + 'Authorization': f'Bearer {self.api_key}', + 'Accept': 'application/json' + } + + # 准备multipart/form-data数据 + data = aiohttp.FormData() + data.add_field('speaker_ids', ','.join(self.speaker_ids)) + data.add_field('file', audio_data, filename='audio.wav', content_type='audio/wav') + + timeout = aiohttp.ClientTimeout(total=10) + + # 网络请求 + async with aiohttp.ClientSession(timeout=timeout) as session: + async with session.post(self.api_url, headers=headers, data=data) as response: + + if response.status == 200: + result = await response.json() + speaker_id = result.get("speaker_id") + score = result.get("score", 0) + total_elapsed_time = time.monotonic() - api_start_time + + logger.bind(tag=TAG).info(f"声纹识别耗时: {total_elapsed_time:.3f}s") + + # 置信度检查 + if score < 0.5: + logger.bind(tag=TAG).warning(f"声纹识别置信度较低: {score:.3f}") + + if speaker_id and speaker_id in self.speaker_map: + result_name = self.speaker_map[speaker_id]["name"] + return result_name + else: + logger.bind(tag=TAG).warning(f"未识别的说话人ID: {speaker_id}") + return "未知说话人" + else: + logger.bind(tag=TAG).error(f"声纹识别API错误: HTTP {response.status}") + return None + + except asyncio.TimeoutError: + elapsed = time.monotonic() - api_start_time + logger.bind(tag=TAG).error(f"声纹识别超时: {elapsed:.3f}s") + return None + except Exception as e: + elapsed = time.monotonic() - api_start_time + logger.bind(tag=TAG).error(f"声纹识别失败: {e}") + return None diff --git a/main/xiaozhi-server/plugins_func/functions/get_time.py b/main/xiaozhi-server/plugins_func/functions/get_time.py index 44732bba..766e19fd 100644 --- a/main/xiaozhi-server/plugins_func/functions/get_time.py +++ b/main/xiaozhi-server/plugins_func/functions/get_time.py @@ -2,59 +2,27 @@ from datetime import datetime import cnlunar from plugins_func.register import register_function, ToolType, ActionResponse, Action -# 添加星期映射字典 -WEEKDAY_MAP = { - "Monday": "星期一", - "Tuesday": "星期二", - "Wednesday": "星期三", - "Thursday": "星期四", - "Friday": "星期五", - "Saturday": "星期六", - "Sunday": "星期日", -} - -get_time_function_desc = { - "type": "function", - "function": { - "name": "get_time", - "description": "获取今天日期或者当前时间信息", - "parameters": {"type": "object", "properties": {}, "required": []}, - }, -} - - -@register_function("get_time", get_time_function_desc, ToolType.WAIT) -def get_time(): - """ - 获取当前的日期时间信息 - """ - now = datetime.now() - current_time = now.strftime("%H:%M:%S") - current_date = now.strftime("%Y-%m-%d") - current_weekday = WEEKDAY_MAP[now.strftime("%A")] - response_text = ( - f"当前日期: {current_date},当前时间: {current_time}, {current_weekday}" - ) - - return ActionResponse(Action.REQLLM, response_text, None) - - get_lunar_function_desc = { "type": "function", "function": { "name": "get_lunar", "description": ( - "用于获取今天的阴历/农历和黄历信息。" + "用于具体日期的阴历/农历和黄历信息。" "用户可以指定查询内容,如:阴历日期、天干地支、节气、生肖、星座、八字、宜忌等。" "如果没有指定查询内容,则默认查询干支年和农历日期。" + "对于'今天农历是多少'、'今天农历日期'这样的基本查询,请直接使用context中的信息,不要调用此工具。" ), "parameters": { "type": "object", "properties": { + "date": { + "type": "string", + "description": "要查询的日期,格式为YYYY-MM-DD,例如2024-01-01。如果不提供,则使用当前日期", + }, "query": { "type": "string", "description": "要查询的内容,例如阴历日期、天干地支、节日、节气、生肖、星座、八字、宜忌等", - } + }, }, "required": [], }, @@ -63,23 +31,41 @@ get_lunar_function_desc = { @register_function("get_lunar", get_lunar_function_desc, ToolType.WAIT) -def get_lunar(query=None): +def get_lunar(date=None, query=None): """ 用于获取当前的阴历/农历,和天干地支、节气、生肖、星座、八字、宜忌等黄历信息 """ - now = datetime.now() - current_time = now.strftime("%H:%M:%S") + from core.utils.cache.manager import cache_manager, CacheType + + # 如果提供了日期参数,则使用指定日期;否则使用当前日期 + if date: + try: + now = datetime.strptime(date, "%Y-%m-%d") + except ValueError: + return ActionResponse( + Action.REQLLM, + f"日期格式错误,请使用YYYY-MM-DD格式,例如:2024-01-01", + None, + ) + else: + now = datetime.now() + current_date = now.strftime("%Y-%m-%d") - current_weekday = WEEKDAY_MAP[now.strftime("%A")] # 如果 query 为 None,则使用默认文本 if query is None: query = "默认查询干支年和农历日期" + + # 尝试从缓存获取农历信息 + lunar_cache_key = f"lunar_info_{current_date}" + cached_lunar_info = cache_manager.get(CacheType.LUNAR, lunar_cache_key) + if cached_lunar_info: + return ActionResponse(Action.REQLLM, cached_lunar_info, None) + response_text = f"根据以下信息回应用户的查询请求,并提供与{query}相关的信息:\n" lunar = cnlunar.Lunar(now, godType="8char") response_text += ( - f"当前公历日期: {current_date},当前时间: {current_time},{current_weekday}\n" "农历信息:\n" "%s年%s%s\n" % (lunar.lunarYearCn, lunar.lunarMonthCn[:-1], lunar.lunarDayCn) + "干支: %s年 %s月 %s日\n" % (lunar.year8Char, lunar.month8Char, lunar.day8Char) @@ -135,4 +121,7 @@ def get_lunar(query=None): + "(默认返回干支年和农历日期;仅在要求查询宜忌信息时才返回本日宜忌)" ) + # 缓存农历信息 + cache_manager.set(CacheType.LUNAR, lunar_cache_key, response_text) + return ActionResponse(Action.REQLLM, response_text, None) diff --git a/main/xiaozhi-server/plugins_func/functions/get_weather.py b/main/xiaozhi-server/plugins_func/functions/get_weather.py index 75c15c7b..a3af6c03 100644 --- a/main/xiaozhi-server/plugins_func/functions/get_weather.py +++ b/main/xiaozhi-server/plugins_func/functions/get_weather.py @@ -151,20 +151,44 @@ def parse_weather_info(soup): @register_function("get_weather", GET_WEATHER_FUNCTION_DESC, ToolType.SYSTEM_CTL) def get_weather(conn, location: str = None, lang: str = "zh_CN"): - api_host = conn.config["plugins"]["get_weather"].get("api_host", "mj7p3y7naa.re.qweatherapi.com") - api_key = conn.config["plugins"]["get_weather"].get("api_key", "a861d0d5e7bf4ee1a83d9a9e4f96d4da") + from core.utils.cache.manager import cache_manager, CacheType + + api_host = conn.config["plugins"]["get_weather"].get( + "api_host", "mj7p3y7naa.re.qweatherapi.com" + ) + api_key = conn.config["plugins"]["get_weather"].get( + "api_key", "a861d0d5e7bf4ee1a83d9a9e4f96d4da" + ) default_location = conn.config["plugins"]["get_weather"]["default_location"] client_ip = conn.client_ip + # 优先使用用户提供的location参数 if not location: # 通过客户端IP解析城市 if client_ip: - # 动态解析IP对应的城市信息 - ip_info = get_ip_info(client_ip, logger) - location = ip_info.get("city") if ip_info and "city" in ip_info else None + # 先从缓存获取IP对应的城市信息 + cached_ip_info = cache_manager.get(CacheType.IP_INFO, client_ip) + if cached_ip_info: + location = cached_ip_info.get("city") + else: + # 缓存未命中,调用API获取 + ip_info = get_ip_info(client_ip, logger) + if ip_info: + cache_manager.set(CacheType.IP_INFO, client_ip, ip_info) + location = ip_info.get("city") + + if not location: + location = default_location else: - # 若IP解析失败或无IP,使用默认位置 + # 若无IP,使用默认位置 location = default_location + # 尝试从缓存获取完整天气报告 + weather_cache_key = f"full_weather_{location}_{lang}" + cached_weather_report = cache_manager.get(CacheType.WEATHER, weather_cache_key) + if cached_weather_report: + return ActionResponse(Action.REQLLM, cached_weather_report, None) + + # 缓存未命中,获取实时天气数据 city_info = fetch_city_info(location, api_key, api_host) if not city_info: return ActionResponse( @@ -192,4 +216,7 @@ def get_weather(conn, location: str = None, lang: str = "zh_CN"): # 提示语 weather_report += "\n(如需某一天的具体天气,请告诉我日期)" + # 缓存完整的天气报告 + cache_manager.set(CacheType.WEATHER, weather_cache_key, weather_report) + return ActionResponse(Action.REQLLM, weather_report, None) diff --git a/main/xiaozhi-server/requirements.txt b/main/xiaozhi-server/requirements.txt index 4ec8b625..522b129d 100755 --- a/main/xiaozhi-server/requirements.txt +++ b/main/xiaozhi-server/requirements.txt @@ -33,4 +33,5 @@ markitdown==0.1.1 mcp-proxy==0.8.0 PyJWT==2.8.0 psutil==7.0.0 -portalocker==2.10.1 \ No newline at end of file +portalocker==2.10.1 +Jinja2==3.1.6 \ No newline at end of file