From b8e9aded6b36b9807bcbe7ba1facca9140e58509 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=AC=A3=E5=8D=97=E7=A7=91=E6=8A=80?= Date: Fri, 14 Mar 2025 14:10:47 +0800 Subject: [PATCH 1/3] Manage agent api (#334) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * update:更换新表 * update:优化用户表sql语句 --------- Co-authored-by: hrz <1710360675@qq.com> --- .../modules/device/entity/DeviceEntity.java | 5 +- .../modules/sys/entity/SysUserEntity.java | 2 +- .../resources/db/changelog/202503101631.sql | 2 - .../{001create_sys.sql => 202503141335.sql} | 7 +- .../{202503131429.sql => 202503141346.sql} | 216 ++++++++++-------- .../db/changelog/db.changelog-master.yaml | 15 +- 6 files changed, 126 insertions(+), 121 deletions(-) delete mode 100644 main/manager-api/src/main/resources/db/changelog/202503101631.sql rename main/manager-api/src/main/resources/db/changelog/{001create_sys.sql => 202503141335.sql} (95%) rename main/manager-api/src/main/resources/db/changelog/{202503131429.sql => 202503141346.sql} (65%) diff --git a/main/manager-api/src/main/java/xiaozhi/modules/device/entity/DeviceEntity.java b/main/manager-api/src/main/java/xiaozhi/modules/device/entity/DeviceEntity.java index bc1a7e88..ccbe4683 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/device/entity/DeviceEntity.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/device/entity/DeviceEntity.java @@ -31,11 +31,8 @@ public class DeviceEntity { @Schema(description = "设备别名") private String alias; - @Schema(description = "智能体编码") - private String agentCode; - @Schema(description = "智能体ID") - private Long agentId; + private String agentId; @Schema(description = "固件版本号") private String appVersion; diff --git a/main/manager-api/src/main/java/xiaozhi/modules/sys/entity/SysUserEntity.java b/main/manager-api/src/main/java/xiaozhi/modules/sys/entity/SysUserEntity.java index 764ee199..c519dd12 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/sys/entity/SysUserEntity.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/sys/entity/SysUserEntity.java @@ -36,7 +36,7 @@ public class SysUserEntity extends BaseEntity { * 更新者 */ @TableField(fill = FieldFill.INSERT) - private Long create_date; + private Date createDate; /** * 更新者 */ diff --git a/main/manager-api/src/main/resources/db/changelog/202503101631.sql b/main/manager-api/src/main/resources/db/changelog/202503101631.sql deleted file mode 100644 index 4055cad5..00000000 --- a/main/manager-api/src/main/resources/db/changelog/202503101631.sql +++ /dev/null @@ -1,2 +0,0 @@ --- 给用户表添加一个创建者 -ALTER TABLE sys_user ADD COLUMN creator BIGINT COMMENT '创建者'; \ No newline at end of file diff --git a/main/manager-api/src/main/resources/db/changelog/001create_sys.sql b/main/manager-api/src/main/resources/db/changelog/202503141335.sql similarity index 95% rename from main/manager-api/src/main/resources/db/changelog/001create_sys.sql rename to main/manager-api/src/main/resources/db/changelog/202503141335.sql index c5bb0cdf..cc51f054 100644 --- a/main/manager-api/src/main/resources/db/changelog/001create_sys.sql +++ b/main/manager-api/src/main/resources/db/changelog/202503141335.sql @@ -13,10 +13,10 @@ CREATE TABLE sys_user ( status tinyint COMMENT '状态 0:停用 1:正常', create_date datetime COMMENT '创建时间', updater bigint COMMENT '更新者', + creator bigint COMMENT '创建者', update_date datetime COMMENT '更新时间', primary key (id), - unique key uk_username (username), - key idx_create_date (create_date) + unique key uk_username (username) ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='系统用户'; -- 系统用户Token @@ -45,8 +45,7 @@ create table sys_params updater bigint COMMENT '更新者', update_date datetime COMMENT '更新时间', primary key (id), - unique key uk_param_code (param_code), - key idx_create_date (create_date) + unique key uk_param_code (param_code) )ENGINE=InnoDB DEFAULT CHARACTER SET utf8mb4 COMMENT='参数管理'; -- 字典类型 diff --git a/main/manager-api/src/main/resources/db/changelog/202503131429.sql b/main/manager-api/src/main/resources/db/changelog/202503141346.sql similarity index 65% rename from main/manager-api/src/main/resources/db/changelog/202503131429.sql rename to main/manager-api/src/main/resources/db/changelog/202503141346.sql index e939b7fe..de14fd71 100644 --- a/main/manager-api/src/main/resources/db/changelog/202503131429.sql +++ b/main/manager-api/src/main/resources/db/changelog/202503141346.sql @@ -1,12 +1,29 @@ +-- 模型供应器表 +DROP TABLE IF EXISTS `ai_model_provider`; +CREATE TABLE `ai_model_provider` ( + `id` VARCHAR(32) NOT NULL COMMENT '主键', + `model_type` VARCHAR(20) COMMENT '模型类型(Memory/ASR/VAD/LLM/TTS)', + `provider_code` VARCHAR(50) COMMENT '供应器类型', + `name` VARCHAR(50) COMMENT '供应器名称', + `fields` JSON COMMENT '供应器字段列表(JSON格式)', + `sort` INT UNSIGNED DEFAULT 0 COMMENT '排序', + `creator` BIGINT COMMENT '创建者', + `create_date` DATETIME COMMENT '创建时间', + `updater` BIGINT COMMENT '更新者', + `update_date` DATETIME COMMENT '更新时间', + PRIMARY KEY (`id`), + INDEX `idx_ai_model_provider_model_type` (`model_type`) COMMENT '创建模型类型的索引,用于快速查找特定类型下的所有供应器信息' +) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='模型配置表'; + -- 模型配置表 DROP TABLE IF EXISTS `ai_model_config`; CREATE TABLE `ai_model_config` ( - `id` BIGINT NOT NULL COMMENT '主键', + `id` VARCHAR(32) NOT NULL COMMENT '主键', `model_type` VARCHAR(20) COMMENT '模型类型(Memory/ASR/VAD/LLM/TTS)', `model_code` VARCHAR(50) COMMENT '模型编码(如AliLLM、DoubaoTTS)', `model_name` VARCHAR(50) COMMENT '模型名称', `is_default` TINYINT(1) DEFAULT 0 COMMENT '是否默认配置(0否 1是)', - `is_enabled` TINYINT(1) DEFAULT 0 COMMENT '是否启用(原注释有误,应为是否启用而非是否默认配置)', + `is_enabled` TINYINT(1) DEFAULT 0 COMMENT '是否启用', `config_json` JSON COMMENT '模型配置(JSON格式)', `doc_link` VARCHAR(200) COMMENT '官方文档链接', `remark` VARCHAR(255) COMMENT '备注', @@ -15,14 +32,15 @@ CREATE TABLE `ai_model_config` ( `create_date` DATETIME COMMENT '创建时间', `updater` BIGINT COMMENT '更新者', `update_date` DATETIME COMMENT '更新时间', - PRIMARY KEY (`id`) + PRIMARY KEY (`id`), + INDEX `idx_ai_model_config_model_type` (`model_type`) COMMENT '创建模型类型的索引,用于快速查找特定类型下的所有配置信息' ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='模型配置表'; -- TTS 音色表 DROP TABLE IF EXISTS `ai_tts_voice`; CREATE TABLE `ai_tts_voice` ( - `id` BIGINT NOT NULL COMMENT '主键', - `tts_model_id` BIGINT COMMENT '对应 TTS 模型主键', + `id` VARCHAR(32) NOT NULL COMMENT '主键', + `tts_model_id` VARCHAR(32) COMMENT '对应 TTS 模型主键', `name` VARCHAR(20) COMMENT '音色名称', `tts_voice` VARCHAR(50) COMMENT '音色编码', `languages` VARCHAR(50) COMMENT '语言', @@ -33,99 +51,14 @@ CREATE TABLE `ai_tts_voice` ( `create_date` DATETIME COMMENT '创建时间', `updater` BIGINT COMMENT '更新者', `update_date` DATETIME COMMENT '更新时间', - PRIMARY KEY (`id`) -) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='TTS 音色表'; - --- 对话历史表 -DROP TABLE IF EXISTS `ai_chat_history`; -CREATE TABLE `ai_chat_history` ( - `id` BIGINT NOT NULL COMMENT '对话编号', - `user_id` BIGINT COMMENT '用户编号', - `agent_id` BIGINT DEFAULT NULL COMMENT '聊天角色', - `device_id` BIGINT DEFAULT NULL COMMENT '设备编号(原注释有误,应为设备编号)', - `message_count` INT COMMENT '信息汇总', - `sort` INT UNSIGNED DEFAULT 0 COMMENT '排序', - `creator` BIGINT COMMENT '创建者', - `create_date` DATETIME COMMENT '创建时间', - `updater` BIGINT COMMENT '更新者', - `update_date` DATETIME COMMENT '更新时间', - PRIMARY KEY (`id`) -) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='对话历史表'; - --- 对话信息表 -DROP TABLE IF EXISTS `ai_chat_message`; -CREATE TABLE `ai_chat_message` ( - `id` BIGINT NOT NULL AUTO_INCREMENT COMMENT '对话记录唯一标识', - `user_id` BIGINT COMMENT '用户唯一标识', - `chat_id` VARCHAR(64) COMMENT '对话历史 ID', - `agent_name` VARCHAR(64) COMMENT '智能体名称', - `role` ENUM('user', 'assistant') COMMENT '角色(用户或助理)', - `content` TEXT COMMENT '对话内容', - `embedding` TEXT COMMENT '对话内容的嵌入向量(可选)', - `url` VARCHAR(255) COMMENT '相关音频文件的 URL(可选)', - `prompt_tokens` INT UNSIGNED DEFAULT 0 COMMENT '提示令牌数', - `total_tokens` INT UNSIGNED DEFAULT 0 COMMENT '总令牌数', - `completion_tokens` INT UNSIGNED DEFAULT 0 COMMENT '完成令牌数', - `prompt_ms` INT UNSIGNED DEFAULT 0 COMMENT '提示耗时(毫秒)', - `total_ms` INT UNSIGNED DEFAULT 0 COMMENT '总耗时(毫秒)', - `completion_ms` INT UNSIGNED DEFAULT 0 COMMENT '完成耗时(毫秒)', - `sort` INT UNSIGNED DEFAULT 0 COMMENT '排序', - `creator` BIGINT COMMENT '创建者', - `create_date` DATETIME COMMENT '创建时间', - `updater` BIGINT COMMENT '更新者', - `update_date` DATETIME COMMENT '更新时间', PRIMARY KEY (`id`), - INDEX `idx_user_id_chat_id_role` (`user_id`, `chat_id`) COMMENT '用户 ID、聊天会话 ID 和角色的联合索引,用于快速检索对话记录', - INDEX `idx_created_at` (`create_date`) COMMENT '创建时间的索引,用于按时间排序或检索对话记录' -) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='对话信息表'; - --- 设备信息表 -DROP TABLE IF EXISTS `ai_device`; -CREATE TABLE `ai_device` ( - `id` BIGINT NOT NULL COMMENT '设备唯一标识', - `user_id` BIGINT COMMENT '关联用户 ID', - `mac_address` VARCHAR(50) COMMENT 'MAC 地址', - `last_connected_at` DATETIME COMMENT '最后连接时间', - `auto_update` TINYINT UNSIGNED DEFAULT 0 COMMENT '自动更新开关(0 关闭/1 开启)', - `board` VARCHAR(50) COMMENT '设备硬件型号', - `alias` VARCHAR(64) DEFAULT NULL COMMENT '设备别名', - `agent_code` VARCHAR(36) COMMENT '智能体编码', - `agent_id` BIGINT COMMENT '智能体 ID', - `app_version` VARCHAR(20) COMMENT '固件版本号', - `sort` INT UNSIGNED DEFAULT 0 COMMENT '排序', - `creator` BIGINT COMMENT '创建者', - `create_date` DATETIME COMMENT '创建时间', - `updater` BIGINT COMMENT '更新者', - `update_date` DATETIME COMMENT '更新时间', - PRIMARY KEY (`id`) -) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='设备信息表'; - --- 智能体配置表 -DROP TABLE IF EXISTS `ai_agent`; -CREATE TABLE `ai_agent` ( - `id` BIGINT NOT NULL COMMENT '智能体唯一标识', - `user_id` BIGINT COMMENT '所属用户 ID', - `agent_code` VARCHAR(36) COMMENT '智能体唯一凭证', - `agent_name` VARCHAR(64) COMMENT '智能体名称', - `tts_voice` VARCHAR(64) COMMENT '语音合成标识', - `llm_model` VARCHAR(32) COMMENT '大语言模型标识', - `memory` TEXT COMMENT '历史记忆数据', - `character` TEXT COMMENT '角色设定参数', - `long_memory_switch` TINYINT UNSIGNED DEFAULT 0 COMMENT '长期记忆开关', - `lang_code` VARCHAR(10) COMMENT '语言编码', - `language` VARCHAR(10) COMMENT '交互语种', - `sort` INT UNSIGNED DEFAULT 0 COMMENT '排序权重', - `creator` BIGINT COMMENT '创建者 ID', - `created_at` DATETIME COMMENT '创建时间', - `updater` BIGINT COMMENT '更新者 ID', - `updated_at` DATETIME COMMENT '更新时间', - PRIMARY KEY (`id`) -) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='智能体配置表'; + INDEX `idx_ai_tts_voice_tts_model_id` (`tts_model_id`) COMMENT '创建 TTS 模型主键的索引,用于快速查找对应模型的音色信息' +) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='TTS 音色表'; -- 智能体配置模板表 DROP TABLE IF EXISTS `ai_agent_template`; CREATE TABLE `ai_agent_template` ( - `id` BIGINT NOT NULL COMMENT '智能体唯一标识', + `id` VARCHAR(32) NOT NULL COMMENT '智能体唯一标识', `agent_code` VARCHAR(36) COMMENT '智能体编码', `agent_name` VARCHAR(64) COMMENT '智能体名称', `asr_model_id` VARCHAR(32) COMMENT '语音识别模型标识', @@ -133,9 +66,9 @@ CREATE TABLE `ai_agent_template` ( `llm_model_id` VARCHAR(32) COMMENT '大语言模型标识', `tts_model_id` VARCHAR(32) COMMENT '语音合成模型标识', `tts_voice_id` VARCHAR(32) COMMENT '音色标识', - `memory` TEXT COMMENT '历史记忆数据', - `character` TEXT COMMENT '角色设定参数', - `long_memory_switch` TINYINT UNSIGNED DEFAULT 0 COMMENT '长期记忆开关', + `mem_model_id` VARCHAR(32) COMMENT '记忆模型标识', + `intent_model_id` VARCHAR(32) COMMENT '意图模型标识', + `system_prompt` TEXT COMMENT '角色设定参数', `lang_code` VARCHAR(10) COMMENT '语言编码', `language` VARCHAR(10) COMMENT '交互语种', `sort` INT UNSIGNED DEFAULT 0 COMMENT '排序权重', @@ -146,13 +79,60 @@ CREATE TABLE `ai_agent_template` ( PRIMARY KEY (`id`) ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='智能体配置模板表'; +-- 智能体配置表 +DROP TABLE IF EXISTS `ai_agent`; +CREATE TABLE `ai_agent` ( + `id` VARCHAR(32) NOT NULL COMMENT '智能体唯一标识', + `user_id` BIGINT COMMENT '所属用户 ID', + `agent_code` VARCHAR(36) COMMENT '智能体编码', + `agent_name` VARCHAR(64) COMMENT '智能体名称', + `asr_model_id` VARCHAR(32) COMMENT '语音识别模型标识', + `vad_model_id` VARCHAR(64) COMMENT '语音活动检测标识', + `llm_model_id` VARCHAR(32) COMMENT '大语言模型标识', + `tts_model_id` VARCHAR(32) COMMENT '语音合成模型标识', + `tts_voice_id` VARCHAR(32) COMMENT '音色标识', + `mem_model_id` VARCHAR(32) COMMENT '记忆模型标识', + `intent_model_id` VARCHAR(32) COMMENT '意图模型标识', + `system_prompt` TEXT COMMENT '角色设定参数', + `lang_code` VARCHAR(10) COMMENT '语言编码', + `language` VARCHAR(10) COMMENT '交互语种', + `sort` INT UNSIGNED DEFAULT 0 COMMENT '排序权重', + `creator` BIGINT COMMENT '创建者 ID', + `created_at` DATETIME COMMENT '创建时间', + `updater` BIGINT COMMENT '更新者 ID', + `updated_at` DATETIME COMMENT '更新时间', + PRIMARY KEY (`id`), + INDEX `idx_ai_agent_user_id` (`user_id`) COMMENT '创建用户的索引,用于快速查找用户下的智能体信息' +) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='智能体配置表'; + +-- 设备信息表 +DROP TABLE IF EXISTS `ai_device`; +CREATE TABLE `ai_device` ( + `id` VARCHAR(32) NOT NULL COMMENT '设备唯一标识', + `user_id` BIGINT COMMENT '关联用户 ID', + `mac_address` VARCHAR(50) COMMENT 'MAC 地址', + `last_connected_at` DATETIME COMMENT '最后连接时间', + `auto_update` TINYINT UNSIGNED DEFAULT 0 COMMENT '自动更新开关(0 关闭/1 开启)', + `board` VARCHAR(50) COMMENT '设备硬件型号', + `alias` VARCHAR(64) DEFAULT NULL COMMENT '设备别名', + `agent_id` VARCHAR(32) COMMENT '智能体 ID', + `app_version` VARCHAR(20) COMMENT '固件版本号', + `sort` INT UNSIGNED DEFAULT 0 COMMENT '排序', + `creator` BIGINT COMMENT '创建者', + `create_date` DATETIME COMMENT '创建时间', + `updater` BIGINT COMMENT '更新者', + `update_date` DATETIME COMMENT '更新时间', + PRIMARY KEY (`id`), + INDEX `idx_ai_device_created_at` (`mac_address`) COMMENT '创建mac的索引,用于快速查找设备信息' +) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='设备信息表'; + -- 声纹识别表 DROP TABLE IF EXISTS `ai_voiceprint`; CREATE TABLE `ai_voiceprint` ( - `id` BIGINT NOT NULL AUTO_INCREMENT COMMENT '声纹唯一标识', + `id` VARCHAR(32) NOT NULL COMMENT '声纹唯一标识', `name` VARCHAR(64) COMMENT '声纹名称', `user_id` BIGINT COMMENT '用户 ID(关联用户表)', - `agent_id` BIGINT COMMENT '关联智能体 ID', + `agent_id` VARCHAR(32) COMMENT '关联智能体 ID', `agent_code` VARCHAR(36) COMMENT '关联智能体编码', `agent_name` VARCHAR(36) COMMENT '关联智能体名称', `description` VARCHAR(255) COMMENT '声纹描述', @@ -164,4 +144,42 @@ CREATE TABLE `ai_voiceprint` ( `updater` BIGINT COMMENT '更新者 ID', `updated_at` DATETIME COMMENT '更新时间', PRIMARY KEY (`id`) -) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='声纹识别表'; \ No newline at end of file +) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='声纹识别表'; + +-- 对话历史表 +DROP TABLE IF EXISTS `ai_chat_history`; +CREATE TABLE `ai_chat_history` ( + `id` VARCHAR(32) NOT NULL COMMENT '对话编号', + `user_id` BIGINT COMMENT '用户编号', + `agent_id` VARCHAR(32) DEFAULT NULL COMMENT '聊天角色', + `device_id` VARCHAR(32) DEFAULT NULL COMMENT '设备编号', + `message_count` INT COMMENT '信息汇总', + `creator` BIGINT COMMENT '创建者', + `create_date` DATETIME COMMENT '创建时间', + `updater` BIGINT COMMENT '更新者', + `update_date` DATETIME COMMENT '更新时间', + PRIMARY KEY (`id`) +) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='对话历史表'; + +-- 对话信息表 +DROP TABLE IF EXISTS `ai_chat_message`; +CREATE TABLE `ai_chat_message` ( + `id` VARCHAR(32) NOT NULL COMMENT '对话记录唯一标识', + `user_id` BIGINT COMMENT '用户唯一标识', + `chat_id` VARCHAR(64) COMMENT '对话历史 ID', + `role` ENUM('user', 'assistant') COMMENT '角色(用户或助理)', + `content` TEXT COMMENT '对话内容', + `prompt_tokens` INT UNSIGNED DEFAULT 0 COMMENT '提示令牌数', + `total_tokens` INT UNSIGNED DEFAULT 0 COMMENT '总令牌数', + `completion_tokens` INT UNSIGNED DEFAULT 0 COMMENT '完成令牌数', + `prompt_ms` INT UNSIGNED DEFAULT 0 COMMENT '提示耗时(毫秒)', + `total_ms` INT UNSIGNED DEFAULT 0 COMMENT '总耗时(毫秒)', + `completion_ms` INT UNSIGNED DEFAULT 0 COMMENT '完成耗时(毫秒)', + `creator` BIGINT COMMENT '创建者', + `create_date` DATETIME COMMENT '创建时间', + `updater` BIGINT COMMENT '更新者', + `update_date` DATETIME COMMENT '更新时间', + PRIMARY KEY (`id`), + INDEX `idx_ai_chat_message_user_id_chat_id_role` (`user_id`, `chat_id`) COMMENT '用户 ID、聊天会话 ID 和角色的联合索引,用于快速检索对话记录', + INDEX `idx_ai_chat_message_created_at` (`create_date`) COMMENT '创建时间的索引,用于按时间排序或检索对话记录' +) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='对话信息表'; 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 d1d7a557..89ed7ff5 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 @@ -3,23 +3,16 @@ # 每次对数据表进行改动时,只允许新建新对changeSet,不允许对上一个changeSet配置及文件进行修改 databaseChangeLog: - changeSet: - id: 001create_sys + id: 202503141335 author: John changes: - sqlFile: encoding: utf8 - path: classpath:db/changelog/001create_sys.sql + path: classpath:db/changelog/202503141335.sql - changeSet: - id: 202503101631 - author: zjy - changes: - - sqlFile: - encoding: utf8 - path: classpath:db/changelog/202503101631.sql - - changeSet: - id: 202503131429 + id: 202503141346 author: czc changes: - sqlFile: encoding: utf8 - path: classpath:db/changelog/202503131429.sql \ No newline at end of file + path: classpath:db/changelog/202503141346.sql \ No newline at end of file From fc3f9823098851b308a409638cd763e9e9e7dacc Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=AC=A3=E5=8D=97=E7=A7=91=E6=8A=80?= Date: Fri, 14 Mar 2025 23:48:59 +0800 Subject: [PATCH 2/3] =?UTF-8?q?=E4=BC=98=E5=8C=96=E7=99=BB=E5=BD=95=20toke?= =?UTF-8?q?n=20=E6=A0=A1=E9=AA=8C=20(#345)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * 获得获取所有模型名称等功能 (#343) * 优化登录 token 校验 (#330) * fix manager-api bug * 优化登录注册流程,跑通注册登录、接口 * update package * 优化登录 token 校验: 1、后端新增 /api/v1/user/info 接口,优化 Oauth2Filter.getRequestToken 逻辑; 2、前端 /api/v1/user/login 登录成功后,保存 token 至浏览器本地; 3、前端请求添加本地 token。 --------- Co-authored-by: 欣南科技 * update:分离home页面的组件 --------- Co-authored-by: CGD <3030332422@qq.com> Co-authored-by: zhisheng Co-authored-by: hrz <1710360675@qq.com> --- .../xiaozhi/common/constant/Constant.java | 2 + .../modules/security/config/ShiroConfig.java | 6 + .../security/controller/LoginController.java | 25 +- .../modules/security/oauth2/Oauth2Filter.java | 15 +- .../modules/security/oauth2/Oauth2Realm.java | 10 + .../security/service/SysUserTokenService.java | 6 +- .../service/impl/SysUserTokenServiceImpl.java | 24 +- .../modules/sys/service/SysUserService.java | 2 + .../sys/service/impl/SysUserServiceImpl.java | 7 + main/manager-web/src/App.vue | 3 + main/manager-web/src/apis/httpRequest.js | 10 +- main/manager-web/src/apis/module/user.js | 81 ++- .../src/components/AddDeviceDialog.vue | 75 ++ .../manager-web/src/components/DeviceItem.vue | 83 +++ main/manager-web/src/components/HeaderBar.vue | 122 ++++ .../manager-web/src/components/HelloWorld.vue | 58 -- main/manager-web/src/router/index.js | 26 +- main/manager-web/src/store/index.js | 12 + main/manager-web/src/views/home.vue | 663 ++---------------- main/manager-web/src/views/login.vue | 28 +- main/manager-web/src/views/roleConfig.vue | 264 +++++++ 21 files changed, 834 insertions(+), 688 deletions(-) create mode 100644 main/manager-web/src/components/AddDeviceDialog.vue create mode 100644 main/manager-web/src/components/DeviceItem.vue create mode 100644 main/manager-web/src/components/HeaderBar.vue delete mode 100644 main/manager-web/src/components/HelloWorld.vue create mode 100644 main/manager-web/src/views/roleConfig.vue 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 ffe7ff3c..10135cef 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 @@ -78,6 +78,8 @@ public interface Constant { */ String TOKEN_HEADER = "token"; + String AUTHORIZATION = "Authorization"; + /** * 路径分割符 */ diff --git a/main/manager-api/src/main/java/xiaozhi/modules/security/config/ShiroConfig.java b/main/manager-api/src/main/java/xiaozhi/modules/security/config/ShiroConfig.java index f770c79f..35cde8aa 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/security/config/ShiroConfig.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/security/config/ShiroConfig.java @@ -58,6 +58,12 @@ public class ShiroConfig { filters.put("oauth2", new Oauth2Filter()); shiroFilter.setFilters(filters); + //添加Shiro的内置过滤器 + /*anon:无需认证就可以访问 + authc:必须认证了才能让问 + user:必须拥有,记住我功能,才能访问 + perms:拥有对某个资源的权限才能访问 + role:拥有某个角色权限才能访问*/ Map filterMap = new LinkedHashMap<>(); filterMap.put("/webjars/**", "anon"); filterMap.put("/druid/**", "anon"); diff --git a/main/manager-api/src/main/java/xiaozhi/modules/security/controller/LoginController.java b/main/manager-api/src/main/java/xiaozhi/modules/security/controller/LoginController.java index 85c5952b..af3f451c 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/security/controller/LoginController.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/security/controller/LoginController.java @@ -4,11 +4,16 @@ import io.swagger.v3.oas.annotations.Operation; import io.swagger.v3.oas.annotations.tags.Tag; import jakarta.servlet.http.HttpServletResponse; import lombok.AllArgsConstructor; +import org.apache.commons.lang3.StringUtils; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; import org.springframework.web.bind.annotation.*; import xiaozhi.common.exception.ErrorCode; import xiaozhi.common.exception.RenException; +import xiaozhi.common.page.TokenDTO; import xiaozhi.common.utils.Result; import xiaozhi.common.validator.AssertUtils; +import xiaozhi.modules.security.dao.SysUserTokenDao; import xiaozhi.modules.security.dto.LoginDTO; import xiaozhi.modules.security.password.PasswordUtils; import xiaozhi.modules.security.service.CaptchaService; @@ -31,6 +36,7 @@ public class LoginController { private final SysUserTokenService sysUserTokenService; private final CaptchaService captchaService; + private static final Logger logger = LoggerFactory.getLogger(LoginController.class); @GetMapping("/captcha") @Operation(summary = "验证码") @@ -44,7 +50,7 @@ public class LoginController { @PostMapping("/login") @Operation(summary = "登录") - public Result login( @RequestBody LoginDTO login) { + public Result login(@RequestBody LoginDTO login) { // 验证是否正确输入验证码 boolean validate = captchaService.validate(login.getCaptchaId(), login.getCaptcha()); if (!validate) { @@ -84,4 +90,21 @@ public class LoginController { } + @GetMapping("/info") + @Operation(summary = "用户信息获取") + public Result info(@RequestHeader("Authorization")String authorization) { + logger.info("the authorization:{}", authorization); + + String token; + if (StringUtils.isBlank(authorization) && authorization.contains("Bearer ")) { + throw new RenException(ErrorCode.UNAUTHORIZED); + } + token = authorization.replace("Bearer ", ""); + if (StringUtils.isBlank(token)) { + throw new RenException(ErrorCode.UNAUTHORIZED); + } + SysUserDTO sysUserDTO = sysUserTokenService.getUserByToken(token); + Result result = new Result(); + return result.ok(sysUserDTO); + } } \ No newline at end of file diff --git a/main/manager-api/src/main/java/xiaozhi/modules/security/oauth2/Oauth2Filter.java b/main/manager-api/src/main/java/xiaozhi/modules/security/oauth2/Oauth2Filter.java index 3d8cf0cb..26486b7b 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/security/oauth2/Oauth2Filter.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/security/oauth2/Oauth2Filter.java @@ -11,6 +11,7 @@ import org.apache.shiro.web.filter.authc.AuthenticatingFilter; import org.springframework.web.bind.annotation.RequestMethod; import xiaozhi.common.constant.Constant; import xiaozhi.common.exception.ErrorCode; +import xiaozhi.common.exception.RenException; import xiaozhi.common.utils.HttpContextUtils; import xiaozhi.common.utils.JsonUtils; import xiaozhi.common.utils.Result; @@ -89,12 +90,22 @@ public class Oauth2Filter extends AuthenticatingFilter { * 获取请求的token */ private String getRequestToken(HttpServletRequest httpRequest) { + String token; //从header中获取token - String token = httpRequest.getHeader(Constant.TOKEN_HEADER); + String authorization = httpRequest.getHeader(Constant.AUTHORIZATION); + if (StringUtils.isBlank(authorization) && authorization.contains("Bearer ")) { + throw new RenException(ErrorCode.UNAUTHORIZED); + } + token = authorization.replace("Bearer ", ""); //如果header中不存在token,则从参数中获取token if (StringUtils.isBlank(token)) { - token = httpRequest.getParameter(Constant.TOKEN_HEADER); + authorization = httpRequest.getParameter(Constant.AUTHORIZATION); + + if (StringUtils.isBlank(authorization) && authorization.contains("Bearer ")) { + throw new RenException(ErrorCode.UNAUTHORIZED); + } + token = authorization.replace("Bearer ", ""); } return token; } diff --git a/main/manager-api/src/main/java/xiaozhi/modules/security/oauth2/Oauth2Realm.java b/main/manager-api/src/main/java/xiaozhi/modules/security/oauth2/Oauth2Realm.java index 3f8a3eeb..8c297e17 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/security/oauth2/Oauth2Realm.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/security/oauth2/Oauth2Realm.java @@ -6,12 +6,15 @@ import org.apache.shiro.authz.AuthorizationInfo; import org.apache.shiro.authz.SimpleAuthorizationInfo; import org.apache.shiro.realm.AuthorizingRealm; import org.apache.shiro.subject.PrincipalCollection; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; import org.springframework.context.annotation.Lazy; import org.springframework.stereotype.Component; import xiaozhi.common.exception.ErrorCode; import xiaozhi.common.user.UserDetail; import xiaozhi.common.utils.ConvertUtils; import xiaozhi.common.utils.MessageUtils; +import xiaozhi.modules.security.controller.LoginController; import xiaozhi.modules.security.entity.SysUserTokenEntity; import xiaozhi.modules.security.service.ShiroService; import xiaozhi.modules.sys.entity.SysUserEntity; @@ -30,6 +33,8 @@ public class Oauth2Realm extends AuthorizingRealm { @Resource private ShiroService shiroService; + private static final Logger logger = LoggerFactory.getLogger(Oauth2Realm.class); + @Override public boolean supports(AuthenticationToken token) { return token instanceof Oauth2Token; @@ -82,6 +87,11 @@ public class Oauth2Realm extends AuthorizingRealm { userDetail.setToken(accessToken); //账号锁定 + if (userDetail.getStatus() == null) { + logger.error("账号状态异常,status 不能为空"); + throw new DisabledAccountException(MessageUtils.getMessage(ErrorCode.ACCOUNT_DISABLE)); + } + if (userDetail.getStatus() == 0) { throw new LockedAccountException(MessageUtils.getMessage(ErrorCode.ACCOUNT_LOCK)); } diff --git a/main/manager-api/src/main/java/xiaozhi/modules/security/service/SysUserTokenService.java b/main/manager-api/src/main/java/xiaozhi/modules/security/service/SysUserTokenService.java index 220b18b5..74882dea 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/security/service/SysUserTokenService.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/security/service/SysUserTokenService.java @@ -1,9 +1,11 @@ package xiaozhi.modules.security.service; import xiaozhi.common.page.PageData; +import xiaozhi.common.page.TokenDTO; import xiaozhi.common.service.BaseService; import xiaozhi.common.utils.Result; import xiaozhi.modules.security.entity.SysUserTokenEntity; +import xiaozhi.modules.sys.dto.SysUserDTO; import java.util.Map; @@ -19,7 +21,9 @@ public interface SysUserTokenService extends BaseService { * * @param userId 用户ID */ - Result createToken(Long userId); + Result createToken(Long userId); + + SysUserDTO getUserByToken(String token); /** * 退出 diff --git a/main/manager-api/src/main/java/xiaozhi/modules/security/service/impl/SysUserTokenServiceImpl.java b/main/manager-api/src/main/java/xiaozhi/modules/security/service/impl/SysUserTokenServiceImpl.java index 903ca242..caa24801 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/security/service/impl/SysUserTokenServiceImpl.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/security/service/impl/SysUserTokenServiceImpl.java @@ -1,6 +1,9 @@ package xiaozhi.modules.security.service.impl; import cn.hutool.core.date.DateUtil; +import lombok.AllArgsConstructor; +import xiaozhi.common.exception.ErrorCode; +import xiaozhi.common.exception.RenException; import xiaozhi.common.page.TokenDTO; import xiaozhi.common.service.impl.BaseServiceImpl; import xiaozhi.common.utils.HttpContextUtils; @@ -10,18 +13,23 @@ import xiaozhi.modules.security.entity.SysUserTokenEntity; import xiaozhi.modules.security.oauth2.TokenGenerator; import xiaozhi.modules.security.service.SysUserTokenService; import org.springframework.stereotype.Service; +import xiaozhi.modules.sys.dto.SysUserDTO; +import xiaozhi.modules.sys.service.SysUserService; import java.util.Date; +@AllArgsConstructor @Service public class SysUserTokenServiceImpl extends BaseServiceImpl implements SysUserTokenService { + + private final SysUserService sysUserService; /** * 12小时后过期 */ private final static int EXPIRE = 3600 * 12; @Override - public Result createToken(Long userId) { + public Result createToken(Long userId) { //用户token String token; @@ -70,6 +78,20 @@ public class SysUserTokenServiceImpl extends BaseServiceImpl { SysUserDTO getByUsername(String username); + SysUserDTO getByUserId(Long userId); + void save(SysUserDTO dto); void delete(Long[] ids); diff --git a/main/manager-api/src/main/java/xiaozhi/modules/sys/service/impl/SysUserServiceImpl.java b/main/manager-api/src/main/java/xiaozhi/modules/sys/service/impl/SysUserServiceImpl.java index 5c3a10c5..8f576137 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/sys/service/impl/SysUserServiceImpl.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/sys/service/impl/SysUserServiceImpl.java @@ -41,6 +41,13 @@ public class SysUserServiceImpl extends BaseServiceImpl + \ No newline at end of file diff --git a/main/manager-web/src/apis/httpRequest.js b/main/manager-web/src/apis/httpRequest.js index 76acd776..98f47bf7 100755 --- a/main/manager-web/src/apis/httpRequest.js +++ b/main/manager-web/src/apis/httpRequest.js @@ -1,6 +1,7 @@ -import {goToPage, showDanger, showWarning} from '../utils/index' +import {goToPage, showDanger, showWarning, isNotNull} from '../utils/index' import Constant from '../utils/constant' import Fly from 'flyio/dist/npm/fly'; +import store from '../store/index' const fly = new Fly() // 设置超时 @@ -25,7 +26,9 @@ function sendRequest() { _url: '', _responseType: undefined, // 新增响应类型字段 'send'() { - this._header.token = localStorage.getItem(Constant.STORAGE_KEY.TOKEN) + if(isNotNull(store.getters.getToken)){ + this._header.Authorization = 'Bearer ' + (JSON.parse(store.getters.getToken)).token + } // 打印请求信息 fly.request(this._url, this._data, { @@ -43,7 +46,7 @@ function sendRequest() { } }).catch((res) => { // 打印失败响应 - console.log(res) + console.log('catch', res) httpHandlerError(res, this._failCallback) }) return this @@ -97,6 +100,7 @@ function sendRequest() { */ // 在错误处理函数中添加日志 function httpHandlerError(info, callBack) { + console.log('httpHandlerError', info) /** 请求成功,退出该函数 可以根据项目需求来判断是否请求成功。这里判断的是status为200的时候是成功 */ let networkError = false diff --git a/main/manager-web/src/apis/module/user.js b/main/manager-web/src/apis/module/user.js index 94de7b3f..444d5d76 100755 --- a/main/manager-web/src/apis/module/user.js +++ b/main/manager-web/src/apis/module/user.js @@ -5,7 +5,8 @@ import {getServiceUrl} from '../api' export default { // 登录 login(loginForm, callback) { - RequestService.sendRequest().url(`${getServiceUrl()}/api/v1/user/login`).method('POST') + RequestService.sendRequest().url(`${getServiceUrl()}/api/v1/user/login`) + .method('POST') .data(loginForm) .success((res) => { RequestService.clearRequestTime() @@ -19,7 +20,8 @@ export default { }, // 获取用户信息 getUserInfo(callback) { - RequestService.sendRequest().url(`${getServiceUrl()}/api/v1/user/info`).method('GET') + RequestService.sendRequest().url(`${getServiceUrl()}/api/v1/user/info`) + .method('GET') .success((res) => { RequestService.clearRequestTime() callback(res) @@ -32,7 +34,8 @@ export default { }, // 获取设备信息 getHomeList(callback) { - RequestService.sendRequest().url(`${getServiceUrl()}/api/v1/user/device/bind`).method('GET') + RequestService.sendRequest().url(`${getServiceUrl()}/api/v1/user/device/bind`) + .method('GET') .success((res) => { RequestService.clearRequestTime() callback(res) @@ -76,14 +79,14 @@ export default { }, // 获取验证码 getCaptcha(uuid, callback) { - + RequestService.sendRequest() .url(`${getServiceUrl()}/api/v1/user/captcha?uuid=${uuid}`) .method('GET') .type('blob') .header({ 'Content-Type': 'image/gif', - 'Pragma': 'No-cache', + 'Pragma': 'No-cache', 'Cache-Control': 'no-cache' }) .success((res) => { @@ -91,7 +94,7 @@ export default { callback(res); }) .fail((err) => { // 添加错误参数 - + }).send() }, // 注册账号 @@ -105,4 +108,70 @@ export default { .fail(() => { }).send() }, + + // 保存设备配置 + saveDeviceConfig(device_id, configData, callback) { + RequestService.sendRequest() + .url(`${getServiceUrl()}/api/v1/user/configDevice/${device_id}`) + .method('PUT') + .data(configData) + .success((res) => { + RequestService.clearRequestTime(); + callback(res); + }) + .fail((err) => { + console.error('保存配置失败:', err); + RequestService.reAjaxFun(() => { + this.saveDeviceConfig(device_id, configData, callback); + }); + }).send(); + }, + // 获取设备配置 + getDeviceConfig(device_id, callback) { + RequestService.sendRequest() + .url(`${getServiceUrl()}/api/v1/user/configDevice/${device_id}`) + .method('GET') + .success((res) => { + RequestService.clearRequestTime(); + callback(res); + }) + .fail((err) => { + console.error('获取配置失败:', err); + RequestService.reAjaxFun(() => { + this.getDeviceConfig(device_id, callback); + }); + }).send(); + }, + // 获取所有模型名称 + getModelNames(callback) { + RequestService.sendRequest() + .url(`${getServiceUrl()}/api/v1/models/names`) + .method('GET') + .success((res) => { + RequestService.clearRequestTime(); + callback(res); + }) + .fail(() => { + RequestService.reAjaxFun(() => { + this.getModelNames(callback); + }); + }).send(); + }, + + // 获取模型音色 + getModelVoices(modelName, callback) { + RequestService.sendRequest() + .url(`${getServiceUrl()}/api/v1/models/${modelName}/voices`) + .method('GET') + .success((res) => { + RequestService.clearRequestTime(); + callback(res); + }) + .fail(() => { + RequestService.reAjaxFun(() => { + this.getModelVoices(modelName, callback); + }); + }).send(); + }, + } diff --git a/main/manager-web/src/components/AddDeviceDialog.vue b/main/manager-web/src/components/AddDeviceDialog.vue new file mode 100644 index 00000000..a9eec57c --- /dev/null +++ b/main/manager-web/src/components/AddDeviceDialog.vue @@ -0,0 +1,75 @@ + + + + + \ No newline at end of file diff --git a/main/manager-web/src/components/DeviceItem.vue b/main/manager-web/src/components/DeviceItem.vue new file mode 100644 index 00000000..0ac1f294 --- /dev/null +++ b/main/manager-web/src/components/DeviceItem.vue @@ -0,0 +1,83 @@ + + + + \ No newline at end of file diff --git a/main/manager-web/src/components/HeaderBar.vue b/main/manager-web/src/components/HeaderBar.vue new file mode 100644 index 00000000..ac87a339 --- /dev/null +++ b/main/manager-web/src/components/HeaderBar.vue @@ -0,0 +1,122 @@ + + + + + \ No newline at end of file diff --git a/main/manager-web/src/components/HelloWorld.vue b/main/manager-web/src/components/HelloWorld.vue deleted file mode 100644 index 2589ba46..00000000 --- a/main/manager-web/src/components/HelloWorld.vue +++ /dev/null @@ -1,58 +0,0 @@ - - - - - - diff --git a/main/manager-web/src/router/index.js b/main/manager-web/src/router/index.js index 41f69a21..e4b7f700 100644 --- a/main/manager-web/src/router/index.js +++ b/main/manager-web/src/router/index.js @@ -1,8 +1,5 @@ import Vue from 'vue' import VueRouter from 'vue-router' -import Welcome from '../views/welcome.vue' -import Login from '../views/login.vue' -import Register from '@/views/register.vue' Vue.use(VueRouter) @@ -10,33 +7,36 @@ const routes = [ { path: '/', name: 'welcome', - component: Login + component: function () { + return import('../views/login.vue') + } + }, + { + path: '/role-config', + name: 'RoleConfig', + component: function () { + return import('../views/roleConfig.vue') + } }, { path: '/login', name: 'login', - // route level code-splitting - // this generates a separate chunk (about.[hash].js) for this route - // which is lazy-loaded when the route is visited. component: function () { - return import(/* webpackChunkName: "about" */ '../views/login.vue') + return import('../views/login.vue') } }, { path: '/home', name: 'home', component: function () { - return import(/* webpackChunkName: "about" */ '../views/home.vue') + return import('../views/home.vue') } }, { path: '/register', name: 'Register', - // route level code-splitting - // this generates a separate chunk (about.[hash].js) for this route - // which is lazy-loaded when the route is visited. component: function () { - return import(/* webpackChunkName: "about" */ '../views/register.vue') + return import('../views/register.vue') } }, ] diff --git a/main/manager-web/src/store/index.js b/main/manager-web/src/store/index.js index ceffa8e3..0f810894 100644 --- a/main/manager-web/src/store/index.js +++ b/main/manager-web/src/store/index.js @@ -1,14 +1,26 @@ import Vue from 'vue' import Vuex from 'vuex' +import Constant from '../utils/constant' Vue.use(Vuex) export default new Vuex.Store({ state: { + token: '' }, getters: { + getToken(state) { + if (!state.token) { + state.token = localStorage.getItem('token') + } + return state.token + } }, mutations: { + setToken(state, token) { + state.token = token + localStorage.token = token + } }, actions: { }, diff --git a/main/manager-web/src/views/home.vue b/main/manager-web/src/views/home.vue index 4d0181f5..173bf8f1 100644 --- a/main/manager-web/src/views/home.vue +++ b/main/manager-web/src/views/home.vue @@ -1,346 +1,88 @@ - - - + \ No newline at end of file diff --git a/main/manager-web/src/views/login.vue b/main/manager-web/src/views/login.vue index d52d2c9b..d6083635 100644 --- a/main/manager-web/src/views/login.vue +++ b/main/manager-web/src/views/login.vue @@ -87,18 +87,22 @@ export default { }, methods: { fetchCaptcha() { - this.captchaUuid = getUUID(); + if (this.$store.getters.getToken) { + goToPage('/home') + } else { + this.captchaUuid = getUUID(); - Api.user.getCaptcha(this.captchaUuid, (res) => { - if (res.status === 200) { - const blob = new Blob([res.data], {type: res.data.type}); - this.captchaUrl = URL.createObjectURL(blob); + Api.user.getCaptcha(this.captchaUuid, (res) => { + if (res.status === 200) { + const blob = new Blob([res.data], {type: res.data.type}); + this.captchaUrl = URL.createObjectURL(blob); - } else { - console.error('验证码加载异常:', error); - showDanger('验证码加载失败,点击刷新') - } - }); + } else { + console.error('验证码加载异常:', error); + showDanger('验证码加载失败,点击刷新') + } + }); + } }, async login() { @@ -119,8 +123,12 @@ export default { Api.user.login(this.form, ({data}) => { console.log(data) showSuccess('登陆成功!') + + this.$store.commit('setToken', JSON.stringify(data.data)) + goToPage('/home') }) + setTimeout(() => { this.fetchCaptcha() }, 1000) diff --git a/main/manager-web/src/views/roleConfig.vue b/main/manager-web/src/views/roleConfig.vue new file mode 100644 index 00000000..8c8dba25 --- /dev/null +++ b/main/manager-web/src/views/roleConfig.vue @@ -0,0 +1,264 @@ + + + + + \ No newline at end of file From 50ecee99ccb00a756fd3f49d1cbf7d57edd2a5bc Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=AC=A3=E5=8D=97=E7=A7=91=E6=8A=80?= Date: Sat, 15 Mar 2025 00:19:45 +0800 Subject: [PATCH 3/3] =?UTF-8?q?ASR=E5=8F=A5=E9=A6=96=E4=B8=A2=E5=AD=97?= =?UTF-8?q?=E9=97=AE=E9=A2=98=20(#346)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix: ASR句首丢字问题 (#338) * feat: 支持dify的多轮对话, chat工作流支持持久化conversation_id #296 (#313) * TTS功能的一些修改 (#307) * fix: tts音频转码兼容 * feat: 支持自定义tts接口服务 * feat: tts超时时间作为可配置项 * feat: 自定义tts音频格式兼容 * style: 使命名更符合语义 --------- Co-authored-by: Jad Co-authored-by: yanyige Co-authored-by: Jad --- main/xiaozhi-server/config.yaml | 17 +++++++++ main/xiaozhi-server/core/connection.py | 5 +-- .../core/handle/musicHandler.py | 2 +- .../core/handle/receiveAudioHandle.py | 5 +-- .../core/providers/llm/dify/dify.py | 9 ++++- .../xiaozhi-server/core/providers/tts/base.py | 15 ++++---- .../core/providers/tts/custom.py | 35 +++++++++++++++++++ 7 files changed, 75 insertions(+), 13 deletions(-) create mode 100644 main/xiaozhi-server/core/providers/tts/custom.py diff --git a/main/xiaozhi-server/config.yaml b/main/xiaozhi-server/config.yaml index da98f1b4..6f93e9c1 100644 --- a/main/xiaozhi-server/config.yaml +++ b/main/xiaozhi-server/config.yaml @@ -58,6 +58,8 @@ delete_audio: true # 没有语音输入多久后断开连接(秒),默认2分钟,即120秒 close_connection_no_voice_time: 120 +# TTS请求超时时间(秒) +tts_timeout: 10 CMD_exit: - "退出" @@ -412,6 +414,21 @@ TTS: # 语速范围0.25-4.0 speed: 1 output_file: tmp/ + CustomTTS: + # 自定义的TTS接口服务,请求参数可自定义 + # 要求接口使用GET方式请求,并返回音频文件 + type: custom + url: "http://127.0.0.1:9880/tts" + params: # 自定义请求参数 + # text: "{prompt_text}" # {prompt_text}会被替换为实际的提示词内容 + # speaker: jok老师 + # speed: 1 + # foo: bar + # testabc: 123456 + headers: # 自定义请求头 + # Authorization: Bearer xxxx + format: wav # 接口返回的音频格式 + output_file: tmp/ # 模块测试配置 module_test: test_sentences: # 自定义测试语句 diff --git a/main/xiaozhi-server/core/connection.py b/main/xiaozhi-server/core/connection.py index 0ff1ac0e..85da8c39 100644 --- a/main/xiaozhi-server/core/connection.py +++ b/main/xiaozhi-server/core/connection.py @@ -451,7 +451,8 @@ class ConnectionHandler: opus_datas, text_index, tts_file = [], 0, None try: self.logger.bind(tag=TAG).debug("正在处理TTS任务...") - tts_file, text, text_index = future.result(timeout=10) + tts_timeout = self.config.get("tts_timeout", 10) + tts_file, text, text_index = future.result(timeout=tts_timeout) if text is None or len(text) <= 0: self.logger.bind(tag=TAG).error(f"TTS出错:{text_index}: tts text is empty") elif tts_file is None: @@ -459,7 +460,7 @@ class ConnectionHandler: else: self.logger.bind(tag=TAG).debug(f"TTS生成:文件路径: {tts_file}") if os.path.exists(tts_file): - opus_datas, duration = self.tts.wav_to_opus_data(tts_file) + opus_datas, duration = self.tts.audio_to_opus_data(tts_file) else: self.logger.bind(tag=TAG).error(f"TTS出错:文件不存在{tts_file}") except TimeoutError: diff --git a/main/xiaozhi-server/core/handle/musicHandler.py b/main/xiaozhi-server/core/handle/musicHandler.py index 022d6dc4..b650b847 100644 --- a/main/xiaozhi-server/core/handle/musicHandler.py +++ b/main/xiaozhi-server/core/handle/musicHandler.py @@ -131,7 +131,7 @@ class MusicHandler: if music_path.endswith(".p3"): opus_packets, duration = p3.decode_opus_from_file(music_path) else: - opus_packets, duration = conn.tts.wav_to_opus_data(music_path) + opus_packets, duration = conn.tts.audio_to_opus_data(music_path) conn.audio_play_queue.put((opus_packets, selected_music, 0)) except Exception as e: diff --git a/main/xiaozhi-server/core/handle/receiveAudioHandle.py b/main/xiaozhi-server/core/handle/receiveAudioHandle.py index bf6d4c36..c5b4135c 100644 --- a/main/xiaozhi-server/core/handle/receiveAudioHandle.py +++ b/main/xiaozhi-server/core/handle/receiveAudioHandle.py @@ -20,7 +20,8 @@ async def handleAudioMessage(conn, audio): # 如果本次没有声音,本段也没声音,就把声音丢弃了 if have_voice == False and conn.client_have_voice == False: await no_voice_close_connect(conn) - conn.asr_audio.clear() + conn.asr_audio.append(audio) + conn.asr_audio = conn.asr_audio[-5:] # 保留最新的5帧音频内容,解决ASR句首丢字问题 return conn.client_no_voice_last_time = 0.0 conn.asr_audio.append(audio) @@ -29,7 +30,7 @@ async def handleAudioMessage(conn, audio): conn.client_abort = False conn.asr_server_receive = False # 音频太短了,无法识别 - if len(conn.asr_audio) < 3: + if len(conn.asr_audio) < 10: conn.asr_server_receive = True else: text, file_path = await conn.asr.speech_to_text(conn.asr_audio, conn.session_id) diff --git a/main/xiaozhi-server/core/providers/llm/dify/dify.py b/main/xiaozhi-server/core/providers/llm/dify/dify.py index 0e9c821a..f3c62639 100644 --- a/main/xiaozhi-server/core/providers/llm/dify/dify.py +++ b/main/xiaozhi-server/core/providers/llm/dify/dify.py @@ -10,11 +10,13 @@ class LLMProvider(LLMProviderBase): def __init__(self, config): self.api_key = config["api_key"] self.base_url = config.get("base_url", "https://api.dify.ai/v1").rstrip('/') + self.session_conversation_map = {} # 存储session_id和conversation_id的映射 def response(self, session_id, dialogue): try: # 取最后一条用户消息 last_msg = next(m for m in reversed(dialogue) if m["role"] == "user") + conversation_id = self.session_conversation_map.get(session_id) # 发起流式请求 with requests.post( @@ -24,13 +26,18 @@ class LLMProvider(LLMProviderBase): "query": last_msg["content"], "response_mode": "streaming", "user": session_id, - "inputs": {} + "inputs": {}, + "conversation_id": conversation_id }, stream=True ) as r: for line in r.iter_lines(): if line.startswith(b'data: '): event = json.loads(line[6:]) + # 如果没有找到conversation_id,则获取此次conversation_id + if not conversation_id: + conversation_id = event.get('conversation_id') + self.session_conversation_map[session_id] = conversation_id # 更新映射 if event.get('answer'): yield event['answer'] diff --git a/main/xiaozhi-server/core/providers/tts/base.py b/main/xiaozhi-server/core/providers/tts/base.py index 84737608..e6f013ca 100644 --- a/main/xiaozhi-server/core/providers/tts/base.py +++ b/main/xiaozhi-server/core/providers/tts/base.py @@ -41,19 +41,20 @@ class TTSProviderBase(ABC): async def text_to_speak(self, text, output_file): pass - def wav_to_opus_data(self, wav_file_path): - # 使用pydub加载PCM文件 + def audio_to_opus_data(self, audio_file_path): + """音频文件转换为Opus编码""" # 获取文件后缀名 - file_type = os.path.splitext(wav_file_path)[1] + file_type = os.path.splitext(audio_file_path)[1] if file_type: file_type = file_type.lstrip('.') - audio = AudioSegment.from_file(wav_file_path, format=file_type) + audio = AudioSegment.from_file(audio_file_path, format=file_type) + # 转换为单声道/16kHz采样率/16位小端编码(确保与编码器匹配) + audio = audio.set_channels(1).set_frame_rate(16000).set_sample_width(2) + + # 音频时长(秒) duration = len(audio) / 1000.0 - # 转换为单声道和16kHz采样率(确保与编码器匹配) - audio = audio.set_channels(1).set_frame_rate(16000) - # 获取原始PCM数据(16位小端) raw_data = audio.raw_data diff --git a/main/xiaozhi-server/core/providers/tts/custom.py b/main/xiaozhi-server/core/providers/tts/custom.py new file mode 100644 index 00000000..d5447878 --- /dev/null +++ b/main/xiaozhi-server/core/providers/tts/custom.py @@ -0,0 +1,35 @@ +import os +import uuid +import requests +from config.logger import setup_logging +from datetime import datetime +from core.providers.tts.base import TTSProviderBase + +TAG = __name__ +logger = setup_logging() + +class TTSProvider(TTSProviderBase): + def __init__(self, config, delete_audio_file): + super().__init__(config, delete_audio_file) + self.url = config.get("url") + self.headers = config.get("headers", {}) + self.params = config.get("params") + self.format = config.get("format", "wav") + self.output_file = config.get("output_file", "tmp/") + + def generate_filename(self): + return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}.{self.format}") + + async def text_to_speak(self, text, output_file): + request_params = {} + for k, v in self.params.items(): + if isinstance(v, str) and "{prompt_text}" in v: + v = v.replace("{prompt_text}", text) + request_params[k] = v + + resp = requests.get(self.url, params=request_params, headers=self.headers) + if resp.status_code == 200: + with open(output_file, "wb") as file: + file.write(resp.content) + else: + logger.bind(tag=TAG).error(f"Custom TTS请求失败: {resp.status_code} - {resp.text}")