diff --git a/.github/ISSUE_TEMPLATE/bug_report.md b/.github/ISSUE_TEMPLATE/bug_report.md new file mode 100644 index 00000000..2d6c173f --- /dev/null +++ b/.github/ISSUE_TEMPLATE/bug_report.md @@ -0,0 +1,34 @@ +--- +name: Bug 报告(Bug Report) +about: 反馈项目中的缺陷或问题 +title: "[Bug] 简短描述问题" +labels: bug +assignees: '' +--- + +## 🐛 问题描述 + + +## 🔍 复现步骤 + +1. 打开 '...' +2. 点击 '...' +3. 滚动到 '...' +4. 看到错误 + +## 🤔 预期行为 + + +## 😯 截图 + + +## 🖥️ 环境信息 +- 操作系统: [例如 Windows 10] +- 浏览器: [例如 Chrome 89] +- 项目版本: [例如 1.0.0] +- Python版本: [例如 3.9.13] +- Jdk版本:[例如 java 21.0.5 2024-10-15 LTS] +- Nodejs版本:[例如 v20.14.0] + +## 📋 其他信息 + diff --git a/.github/ISSUE_TEMPLATE/code_improvement.md b/.github/ISSUE_TEMPLATE/code_improvement.md new file mode 100644 index 00000000..0ec0b145 --- /dev/null +++ b/.github/ISSUE_TEMPLATE/code_improvement.md @@ -0,0 +1,19 @@ +--- +name: 代码优化建议(Code Improvement) +about: 提出对现有代码的优化或改进建议 +title: "[Improvement] 简短描述改进内容" +labels: refactor +assignees: '' +--- + +## 💡 改进描述 + + +## 🌟 改进建议 + + +## 🛠️ 相关代码 + + +## 📋 其他信息 + diff --git a/.github/ISSUE_TEMPLATE/documentation_improvement.md b/.github/ISSUE_TEMPLATE/documentation_improvement.md new file mode 100644 index 00000000..9cdc10e0 --- /dev/null +++ b/.github/ISSUE_TEMPLATE/documentation_improvement.md @@ -0,0 +1,16 @@ +--- +name: 文档改进建议(Documentation Improvement) +about: 提出对项目文档的改进或补充建议 +title: "[Docs] 简短描述改进内容" +labels: documentation +assignees: '' +--- + +## 📚 改进描述 + + +## ✨ 改进建议 + + +## 📋 其他信息 + diff --git a/.github/ISSUE_TEMPLATE/feature_request.md b/.github/ISSUE_TEMPLATE/feature_request.md new file mode 100644 index 00000000..df1c029d --- /dev/null +++ b/.github/ISSUE_TEMPLATE/feature_request.md @@ -0,0 +1,19 @@ +--- +name: 功能请求(Feature Request) +about: 提出新的功能或改进建议 +title: "[Feature] 简短描述功能" +labels: enhancement +assignees: '' +--- + +## 🚀 需求描述 + + +## 🎯 解决方案 + + +## 📝 备选方案 + + +## 📋 其他信息 + diff --git a/README.md b/README.md index 652b97bd..cdf31847 100644 --- a/README.md +++ b/README.md @@ -57,6 +57,13 @@ + + + + 成本最低配置 + + + @@ -74,9 +81,16 @@ - + - 成本最低配置 + 天气插件 + + + + + + + IOT指令控制设备 @@ -92,9 +106,11 @@ - **硬件**:一套兼容 `xiaozhi-esp32` 的硬件设备(具体型号请参考 [此处](https://rcnv1t9vps13.feishu.cn/wiki/DdgIw4BUgivWDPkhMj1cGIYCnRf))。 -- **电脑或服务器**:至少 4 核 CPU、8G 内存的电脑。 +- **电脑或服务器**:建议 4 核 CPU、8G 内存的电脑。如果开启ASR也使用API,可运行在2核CPU、2G内存的服务器中。 - **固件编译**:请将本后端服务的接口地址更新至 `xiaozhi-esp32` 项目中,再重新编译`xiaozhi-esp32`固件并烧录到设备上。 +![图片](docs/images/deploy.png) + 如果你没有esp32相关的硬件设备,但是非常想体验该项目,可以使用以下的项目让你的电脑、手机模拟成esp32设备。 - [小智安卓端](https://github.com/TOM88812/xiaozhi-android-client) @@ -141,12 +157,15 @@ server: 支持 EdgeTTS(默认)、火山引擎豆包 TTS 等多种 TTS 接口,满足语音合成需求。 - **记忆功能** 支持超长记忆、本地总结记忆、无记忆三种模式,满足不同场景需求。 +- **IOT功能** + 支持管理注册设备IOT功能,支持基于对话上下文语境下的智能物联网控制。 ### 正在开发 🚧 - 多种心情模式 - 智控台webui -- iot功能 + +想了解具体开发进度,[请点击这里](https://github.com/users/xinnan-tech/projects/3) ![图片](docs/images/webui.png) --- @@ -210,10 +229,10 @@ server: ### Memory 记忆存储 -| 类型 | 平台名称 | 使用方式 | 收费模式 | 备注 | -|:------:|:---------------:|:----:|:--------:|:--:| -| Memory | mem0ai | 接口调用 | 100次/月额度 | | -| Memory | mem_local_short | 本地总结 | 免费 | | +| 类型 | 平台名称 | 使用方式 | 收费模式 | 备注 | +|:------:|:---------------:|:----:|:---------:|:--:| +| Memory | mem0ai | 接口调用 | 1000次/月额度 | | +| Memory | mem_local_short | 本地总结 | 免费 | | --- diff --git a/docker-setup.sh b/docker-setup.sh new file mode 100755 index 00000000..04ea547a --- /dev/null +++ b/docker-setup.sh @@ -0,0 +1,105 @@ +#!/bin/sh +# 本文件是用于一键自动下载本项目所需文件,自动创建好目录 +# 所需条件(否则无法使用): +# 1、请确保你的环境可以正常访问 GitHub 否则无法下载脚本 +# +# 检测操作系统类型 +case "$(uname -s)" in + Linux*) OS=Linux;; + Darwin*) OS=Mac;; + CYGWIN*) OS=Windows;; + MINGW*) OS=Windows;; + MSYS*) OS=Windows;; + *) OS=UNKNOWN;; +esac + +# 设置颜色(Windows CMD 不支持,但不影响使用) +if [ "$OS" = "Windows" ]; then + GREEN="" + RED="" + NC="" +else + GREEN='\033[0;32m' + RED='\033[0;31m' + NC='\033[0m' +fi + +echo "${GREEN}开始安装小智服务端...${NC}" + +# 创建必要的目录 +echo "创建目录结构..." +mkdir -p xiaozhi-server/data xiaozhi-server/models/SenseVoiceSmall +cd xiaozhi-server || exit + +# 根据操作系统选择下载命令 +if [ "$OS" = "Windows" ]; then + DOWNLOAD_CMD="curl -L -o" + if ! command -v curl >/dev/null 2>&1; then + DOWNLOAD_CMD="powershell -Command Invoke-WebRequest -Uri" + DOWNLOAD_CMD_SUFFIX="-OutFile" + fi +else + if command -v curl >/dev/null 2>&1; then + DOWNLOAD_CMD="curl -L -o" + elif command -v wget >/dev/null 2>&1; then + DOWNLOAD_CMD="wget -O" + else + echo "${RED}错误: 需要安装 curl 或 wget${NC}" + exit 1 + fi +fi + +# 下载语音识别模型 +echo "下载语音识别模型..." +if [ "$DOWNLOAD_CMD" = "powershell -Command Invoke-WebRequest -Uri" ]; then + $DOWNLOAD_CMD "https://modelscope.cn/models/iic/SenseVoiceSmall/resolve/master/model.pt" $DOWNLOAD_CMD_SUFFIX "models/SenseVoiceSmall/model.pt" +else + $DOWNLOAD_CMD "models/SenseVoiceSmall/model.pt" "https://modelscope.cn/models/iic/SenseVoiceSmall/resolve/master/model.pt" +fi + +if [ $? -ne 0 ]; then + echo "${RED}模型下载失败。请手动从以下地址下载:${NC}" + echo "1. https://modelscope.cn/models/iic/SenseVoiceSmall/resolve/master/model.pt" + echo "2. 百度网盘: https://pan.baidu.com/share/init?surl=QlgM58FHhYv1tFnUT_A8Sg (提取码: qvna)" + echo "下载后请将文件放置在 models/SenseVoiceSmall/model.pt" +fi + +# 下载配置文件 +echo "下载配置文件..." +if [ "$DOWNLOAD_CMD" = "powershell -Command Invoke-WebRequest -Uri" ]; then + $DOWNLOAD_CMD "https://raw.githubusercontent.com/xinnan-tech/xiaozhi-esp32-server/main/main/xiaozhi-server/docker-compose.yml" $DOWNLOAD_CMD_SUFFIX "docker-compose.yml" + $DOWNLOAD_CMD "https://raw.githubusercontent.com/xinnan-tech/xiaozhi-esp32-server/main/main/xiaozhi-server/config.yaml" $DOWNLOAD_CMD_SUFFIX "data/.config.yaml" +else + $DOWNLOAD_CMD "docker-compose.yml" "https://raw.githubusercontent.com/xinnan-tech/xiaozhi-esp32-server/main/main/xiaozhi-server/docker-compose.yml" + $DOWNLOAD_CMD "data/.config.yaml" "https://raw.githubusercontent.com/xinnan-tech/xiaozhi-esp32-server/main/main/xiaozhi-server/config.yaml" +fi + +# 检查文件是否存在 +echo "检查文件完整性..." +FILES_TO_CHECK="docker-compose.yml data/.config.yaml models/SenseVoiceSmall/model.pt" +ALL_FILES_EXIST=true + +for FILE in $FILES_TO_CHECK; do + if [ ! -f "$FILE" ]; then + echo "${RED}错误: $FILE 不存在${NC}" + ALL_FILES_EXIST=false + fi +done + +if [ "$ALL_FILES_EXIST" = false ]; then + echo "${RED}某些文件下载失败,请检查上述错误信息并手动下载缺失的文件。${NC}" + exit 1 +fi + +echo "${GREEN}文件下载完成!${NC}" +echo "请编辑 data/.config.yaml 文件配置你的API密钥。" +echo "配置完成后,运行以下命令启动服务:" +echo "${GREEN}docker-compose up -d${NC}" +echo "查看日志请运行:" +echo "${GREEN}docker logs -f xiaozhi-esp32-server${NC}" + +# 提示用户编辑配置文件 +echo "\n${RED}重要提示:${NC}" +echo "1. 请确保编辑 data/.config.yaml 文件,配置必要的API密钥" +echo "2. 特别是 ChatGLM 和 mem0ai 的密钥必须配置" +echo "3. 配置完成后再启动 docker 服务" diff --git a/docs/Deployment.md b/docs/Deployment.md index 94b6e363..8765d26f 100644 --- a/docs/Deployment.md +++ b/docs/Deployment.md @@ -1,3 +1,5 @@ +# 部署方案参考 +![图片](images/deploy.png) # 方式一:docker快速部署 docker镜像已支持x86架构、arm64架构的CPU,支持在国产操作系统上运行。 @@ -6,7 +8,45 @@ docker镜像已支持x86架构、arm64架构的CPU,支持在国产操作系统 如果您的电脑还没安装docker,可以按照这里的教程安装:[docker安装](https://www.runoob.com/docker/ubuntu-docker-install.html) -## 2. 创建目录 +如果你已经安装好docker,你可以[1.1使用懒人脚本](#11-懒人脚本)自动帮你下载所需的文件和配置文件,你可以使用docker[1.2手动部署](#12-手动部署)。 + +### 1.1 懒人脚本 + +你可以使用以下命令一键下载并执行部署脚本: +请确保你的环境可以正常访问 GitHub 否则无法下载脚本。 +```bash +curl -L -o docker-setup.sh https://raw.githubusercontent.com/xinnan-tech/xiaozhi-esp32-server/main/docker-setup.sh +``` + +如果您的电脑是windows系统,请使用使用 Git Bash、WSL、PowerShell 或 CMD 运行以下命令: +```bash +# Git Bash 或 WSL +sh docker-setup.sh +# PowerShell 或 CMD +.\docker-setup.sh +``` + +如果您的电脑是linux 或者 macos 系统,请使用终端运行以下命令: +```bash +chmod +x docker-setup.sh +./docker-setup.sh +``` + +脚本会自动完成以下操作: +> 1. 创建必要的目录结构 +> 2. 下载语音识别模型 +> 3. 下载配置文件 +> 4. 检查文件完整性 +> +> 执行完成后,请按照提示配置 API 密钥。 + +当你一切顺利完成以上操作后,继续操作[配置项目文件](#3-配置项目文件) + +### 1.2 手动部署 + +如果懒人脚本无法正常运行,请按本章节1.2进行手动部署。 + +#### 1.2.1 创建目录 安装完后,你需要为这个项目找一个安放配置文件的目录,例如我们可以新建一个文件夹叫`xiaozhi-server`。 @@ -21,14 +61,18 @@ xiaozhi-server ├─ SenseVoiceSmall ``` -## 4. 下载语音识别模型文件 +#### 1.2.2 下载语音识别模型文件 你需要下载语音识别的模型文件,因为本项目的默认语音识别用的是本地离线语音识别方案。可通过这个方式下载 [跳转到下载语音识别模型文件](#模型文件) 下载完后,回到本教程。 -## 3. 下载docker-compose.yaml +#### 1.2.3 下载配置文件 + +你需要下载两个配置文件:`docker-compose.yaml` 和 `config.yaml`。需要从项目仓库下载这两个文件。 + +##### 1.2.3.1 下载 docker-compose.yaml 用浏览器打开[这个链接](../main/xiaozhi-server/docker-compose.yml)。 @@ -37,7 +81,7 @@ xiaozhi-server 下载完后,回到本教程继续往下。 -## 3. 下载配置文件 +##### 1.2.3.2 下载 config.yaml 用浏览器打开[这个链接](../main/xiaozhi-server/config.yaml)。 @@ -58,14 +102,14 @@ xiaozhi-server 如果你的文件目录结构也是上面的,就继续往下。如果不是,你就再仔细看看是不是漏操作了什么。 -## 4. 配置项目文件 +## 3. 配置项目文件 接下里,程序还不能直接运行,你需要配置一下,你到底使用的是什么模型。你可以看这个教程: [跳转到配置项目文件](#配置项目) 配置完项目文件后,回到本教程继续往下。 -## 5. 执行docker命令 +## 4. 执行docker命令 打开命令行工具,使用`终端`或`命令行`工具 进入到你的`xiaozhi-server`,执行以下命令 @@ -81,14 +125,14 @@ docker logs -f xiaozhi-esp32-server 这时,你就要留意日志信息,可以根据这个教程,判断是否成功了。[跳转到运行状态确认](#运行状态确认) -## 6.版本升级操作 +## 5. 版本升级操作 如果后期想升级版本,可以这么操作 -1、备份好`data`文件夹中的`.config.yaml`文件,一些关键的配置到时复制到新的`.config.yaml`文件里。 +5.1、备份好`data`文件夹中的`.config.yaml`文件,一些关键的配置到时复制到新的`.config.yaml`文件里。 请注意是对关键密钥逐个复制,不要直接覆盖。因为新的`.config.yaml`文件可能有一些新的配置项,旧的`.config.yaml`文件不一定有。 -2、执行以下命令 +5.2、执行以下命令 ``` docker stop xiaozhi-esp32-server @@ -96,7 +140,7 @@ docker rm xiaozhi-esp32-server docker rmi ghcr.nju.edu.cn/xinnan-tech/xiaozhi-esp32-server:server_latest ``` -3、重新按docker方式部署 +5.3、重新按docker方式部署 # 方式二:借助Docker环境运行部署 @@ -217,20 +261,23 @@ python app.py 如果你的`xiaozhi-server`目录没有`data`,你需要创建`data`目录。 如果你的`data`下面没有`.config.yaml`文件,你可以把源码目录下的`config.yaml`文件复制一份,重命名为`.config.yaml` -修改`xiaozhi-server`下`data`目录下的`.config.yaml`文件,配置本项目必须的两个配置。 +修改`xiaozhi-server`下`data`目录下的`.config.yaml`文件,配置本项目必须的一个配置。 - 默认的LLM使用的是`ChatGLMLLM`,你需要配置密钥,因为他们的模型,虽然有免费的,但是仍要去[官网](https://bigmodel.cn/usercenter/proj-mgmt/apikeys)注册密钥,才能启动。 -- 默认的记忆层`mem0ai`,你需要配置密钥,因为他们的API,虽然有免费额度,但是仍要去[官网](https://app.mem0.ai/dashboard/api-keys)注册密钥,才能启动。 配置说明:这里是各个功能使用的默认组件,例如LLM默认使用`ChatGLMLLM`模型。如果需要切换模型,就是改对应的名称。 本项目的默认配置仅是成本最低配置(`glm-4-flash`和`EdgeTTS`都是免费的),如果需要更优的更快的搭配,需要自己结合部署环境切换各组件的使用。 ``` selected_module: - ASR: FunASR VAD: SileroVAD + ASR: FunASR LLM: ChatGLMLLM TTS: EdgeTTS + # 默认不开启记忆,如需开启请看配置文件里的描述 + Memory: nomem + # 默认不开启意图识别,如需开启请看配置文件里的描述 + Intent: nointent ``` 比如修改`LLM`使用的组件,就看本项目支持哪些`LLM` API接口,当前支持的是`openai`、`dify`。欢迎验证和支持更多LLM平台的接口。 @@ -249,8 +296,6 @@ LLM: ... ``` -有些服务,比如如果你使用`Dify`、`豆包的TTS`,是需要密钥的,记得在配置文件加上哦! - ## 模型文件 本项目语音识别模型,默认使用`SenseVoiceSmall`模型,进行语音转文字。因为模型较大,需要独立下载,下载后把`model.pt` @@ -293,4 +338,4 @@ LLM: [5、我说话很慢,停顿时小智老是抢话](../README.md#1%E4%B8%BA%E4%BB%80%E4%B9%88%E6%88%91%E8%AF%B4%E7%9A%84%E8%AF%9D%E5%B0%8F%E6%99%BA%E8%AF%86%E5%88%AB%E5%87%BA%E6%9D%A5%E5%BE%88%E5%A4%9A%E9%9F%A9%E6%96%87%E6%97%A5%E6%96%87%E8%8B%B1%E6%96%87) -[6、我想通过小智控制电灯、空调、远程开关机等操作](../README.md#1%E4%B8%BA%E4%BB%80%E4%B9%88%E6%88%91%E8%AF%B4%E7%9A%84%E8%AF%9D%E5%B0%8F%E6%99%BA%E8%AF%86%E5%88%AB%E5%87%BA%E6%9D%A5%E5%BE%88%E5%A4%9A%E9%9F%A9%E6%96%87%E6%97%A5%E6%96%87%E8%8B%B1%E6%96%87) +[6、我想通过小智控制电灯、空调、远程开关机等操作](../README.md#1%E4%B8%BA%E4%BB%80%E4%B9%88%E6%88%91%E8%AF%B4%E7%9A%84%E8%AF%9D%E5%B0%8F%E6%99%BA%E8%AF%86%E5%88%AB%E5%87%BA%E6%9D%A5%E5%BE%88%E5%A4%9A%E9%9F%A9%E6%96%87%E6%97%A5%E6%96%87%E8%8B%B1%E6%96%87) \ No newline at end of file diff --git a/docs/contributor_open_letter.md b/docs/contributor_open_letter.md index 1be3da28..7a547197 100644 --- a/docs/contributor_open_letter.md +++ b/docs/contributor_open_letter.md @@ -2,10 +2,27 @@ "春江水暖鸭先知,正是河豚欲上时!" -亲爱的朋友,今天,我怀着无比真挚的心情,向热爱AI技术与创新的你发出这封公开信。 +亲爱的朋友,我是John,是一名普通公司里的Java程序员,今天,我怀着无比真挚的心情,向热爱AI技术与创新的你发出这封公开信。 -我们的开源项目 **xiaozhi-esp32-server** 正处在发展的阶段,它承载着我们对智能硬件和低成本民用贾维斯解决方案的美好憧憬,也希望借助每一位开发者的智慧,共同开创技术的新局面。 +半年前我看到很多优秀的项目,比如`Dify`、`Chat2DB`等人工智能相关的项目,我在想,我要是能参与这些项目多好,可惜“报国无门,空打十年代码”。 +我是2025年初刷到虾哥团队的视频,我非常好奇他是怎么实现的,我想复刻他们的后端服务,打造一个低成本民用贾维斯。很可惜现在做的作品依然只是一个人工智障,它并发低、没有灵魂,响应很慢,bug很多。 + +虾哥团队是我们学习的对象,我很想拥有像虾哥团队一样智能的小智后端服务。但是我也能理解虾哥不开源的决定。“一花独放不是春,百花齐放春满园”,人工智能遍地开花的时代,也许就在我们这代实现,我们可以用自己的双手,实现低成本民用贾维斯。我个人认为,他能实现的,我们也能实现,只是时间问题而已,我称之为“我们的取经之路”。 + +那么这条取经之路,我们会遇到什么困难?我想应该不少于八十一难。这一路必然会出现各种妖怪,当然也有神仙暗中帮助我们,也有人加入取经队伍。 + +以上内容,如果你觉得好笑。那我也觉得非常的幸运。我能够在你人生3万多天里博你笑五秒,也算是为你做了一次贡献。 + +民用低成本贾维斯这个想法,会失败吗,我不知道,但是我们普通人的一生,这种失败不是很常见吗? + +未来,有一点是可以确定的,就一定会有人完全复刻虾哥团队的功能,实现民用低成本贾维斯。这个项目会是我们吗? + +期待与你携手前行,共创未来。 + +John,2025.3.11,广州 + +# 附 开发贡献指南 ## 项目目标 1. **民用低成本贾维斯解决方案** @@ -14,7 +31,7 @@ ## 加入我们 -我们热忱欢迎志同道合的朋友加入,共同为项目贡献力量。参与方式如下: +我们热忱欢迎志同道合的朋友加入,共同为项目贡献力量。您可在[这个链接](https://github.com/users/xinnan-tech/projects/3)查看我们近期要实现的功能,功能列表中还没指派相关人员处理的,正是急需您的参与。参与方式如下: ### 1、成为普通贡献者 @@ -31,21 +48,3 @@ Fork 项目,提交 PR,由开发者审核后合入主分支。 2. **提交 PR 审核** 功能开发完成后,请在 GitHub 上提交 PR,由其他开发者审核,审核通过后合并入主分支。 - -## 功能点的来源 - -- **创新的点子与功能** - 你的每一个灵感都可能成为项目的突破点,欢迎大胆提出前沿的创意与功能建议。 - -- **代码优化** - 如果你发现代码中存在改进的空间,欢迎提交优化方案,让我们的代码更加高效、易读。 - -- **解决 Issues 问题** - Issues 中有很多问题需要解决,我们一同完善、打造产品。 - -亲爱的开发者们, -每一次提交、每一份建议,都是对我们共同梦想的加持。每一代人有每一代人的使命,让我们共同书写智能生活的新篇章! - -期待与你携手前行,共创未来。 - -John,2025.3.11,广州 \ No newline at end of file diff --git a/docs/images/demo8.png b/docs/images/demo8.png new file mode 100644 index 00000000..affe6ce2 Binary files /dev/null and b/docs/images/demo8.png differ diff --git a/docs/images/demo9.png b/docs/images/demo9.png new file mode 100644 index 00000000..74b6e463 Binary files /dev/null and b/docs/images/demo9.png differ diff --git a/docs/images/deploy.png b/docs/images/deploy.png new file mode 100644 index 00000000..e1db8a48 Binary files /dev/null and b/docs/images/deploy.png differ diff --git a/main/README.md b/main/README.md index 6f330023..df4e57a4 100644 --- a/main/README.md +++ b/main/README.md @@ -18,6 +18,6 @@ xiaozhi-esp32-server # manager-web 、manager-api接口协议 -[manager前后端接口协议](https://app.apifox.com/invite/project?token=H_8qhgfjUeaAL0wybghgU) +[manager前后端接口协议](https://app.apifox.com/invite/project?token=eXg2_tUv85q-gc3ZRowmn) [前端页面设计图](https://codesign.qq.com/app/s/526108506410828) diff --git a/main/manager-api/src/main/java/xiaozhi/common/aspect/DataFilterAspect.java b/main/manager-api/src/main/java/xiaozhi/common/aspect/DataFilterAspect.java deleted file mode 100644 index db75689d..00000000 --- a/main/manager-api/src/main/java/xiaozhi/common/aspect/DataFilterAspect.java +++ /dev/null @@ -1,99 +0,0 @@ -package xiaozhi.common.aspect; - -import cn.hutool.core.collection.CollUtil; -import xiaozhi.common.annotation.DataFilter; -import xiaozhi.common.constant.Constant; -import xiaozhi.common.exception.ErrorCode; -import xiaozhi.common.exception.RenException; -import xiaozhi.common.interceptor.DataScope; -import xiaozhi.common.user.UserDetail; -import xiaozhi.modules.security.user.SecurityUser; -import xiaozhi.modules.sys.enums.SuperAdminEnum; -import org.apache.commons.lang3.StringUtils; -import org.aspectj.lang.JoinPoint; -import org.aspectj.lang.annotation.Aspect; -import org.aspectj.lang.annotation.Before; -import org.aspectj.lang.annotation.Pointcut; -import org.aspectj.lang.reflect.MethodSignature; -import org.springframework.stereotype.Component; - -import java.lang.reflect.Method; -import java.util.List; -import java.util.Map; - -/** - * 数据过滤,切面处理类 - * Copyright (c) 人人开源 All rights reserved. - * Website: https://www.renren.io - */ -@Aspect -@Component -public class DataFilterAspect { - - @Pointcut("@annotation(xiaozhi.common.annotation.DataFilter)") - public void dataFilterCut() { - - } - - @Before("dataFilterCut()") - public void dataFilter(JoinPoint point) { - Object params = point.getArgs()[0]; - if (params != null && params instanceof Map) { - UserDetail user = SecurityUser.getUser(); - - //如果是超级管理员,则不进行数据过滤 - if (user.getSuperAdmin() == SuperAdminEnum.YES.value()) { - return; - } - - try { - //否则进行数据过滤 - Map map = (Map) params; - String sqlFilter = getSqlFilter(user, point); - map.put(Constant.SQL_FILTER, new DataScope(sqlFilter)); - } catch (Exception e) { - - } - - return; - } - - throw new RenException(ErrorCode.DATA_SCOPE_PARAMS_ERROR); - } - - /** - * 获取数据过滤的SQL - */ - private String getSqlFilter(UserDetail user, JoinPoint point) throws Exception { - MethodSignature signature = (MethodSignature) point.getSignature(); - Method method = point.getTarget().getClass().getDeclaredMethod(signature.getName(), signature.getParameterTypes()); - DataFilter dataFilter = method.getAnnotation(DataFilter.class); - - //获取表的别名 - String tableAlias = dataFilter.tableAlias(); - if (StringUtils.isNotBlank(tableAlias)) { - tableAlias += "."; - } - - StringBuilder sqlFilter = new StringBuilder(); - sqlFilter.append(" ("); - - //部门ID列表 - List deptIdList = user.getDeptIdList(); - if (CollUtil.isNotEmpty(deptIdList)) { - sqlFilter.append(tableAlias).append(dataFilter.deptId()); - - sqlFilter.append(" in(").append(StringUtils.join(deptIdList, ",")).append(")"); - } - - //查询本人数据 - if (CollUtil.isNotEmpty(deptIdList)) { - sqlFilter.append(" or "); - } - sqlFilter.append(tableAlias).append(dataFilter.userId()).append("=").append(user.getId()); - - sqlFilter.append(")"); - - return sqlFilter.toString(); - } -} \ 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 ffe7ff3c..3738498b 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 @@ -73,10 +73,8 @@ public interface Constant { * 排序方式 */ String ORDER = "order"; - /** - * token header - */ - String TOKEN_HEADER = "token"; + + String AUTHORIZATION = "Authorization"; /** * 路径分割符 diff --git a/main/manager-api/src/main/java/xiaozhi/common/exception/ErrorCode.java b/main/manager-api/src/main/java/xiaozhi/common/exception/ErrorCode.java index e2e5de5d..e5365b3c 100644 --- a/main/manager-api/src/main/java/xiaozhi/common/exception/ErrorCode.java +++ b/main/manager-api/src/main/java/xiaozhi/common/exception/ErrorCode.java @@ -42,4 +42,5 @@ public interface ErrorCode { int PASSWORD_LENGTH_ERROR = 10030; int PASSWORD_WEAK_ERROR = 10031; int DEL_MYSELF_ERROR = 10032; + int DEVICE_CAPTCHA_ERROR = 10033; } diff --git a/main/manager-api/src/main/java/xiaozhi/common/handler/FieldMetaObjectHandler.java b/main/manager-api/src/main/java/xiaozhi/common/handler/FieldMetaObjectHandler.java index fbb75322..20371fd6 100644 --- a/main/manager-api/src/main/java/xiaozhi/common/handler/FieldMetaObjectHandler.java +++ b/main/manager-api/src/main/java/xiaozhi/common/handler/FieldMetaObjectHandler.java @@ -34,9 +34,6 @@ public class FieldMetaObjectHandler implements MetaObjectHandler { //创建时间 strictInsertFill(metaObject, CREATE_DATE, Date.class, date); - //创建者所属部门 - strictInsertFill(metaObject, DEPT_ID, Long.class, user.getDeptId()); - //更新者 strictInsertFill(metaObject, UPDATER, Long.class, user.getId()); //更新时间 diff --git a/main/manager-api/src/main/java/xiaozhi/common/redis/RedisKeys.java b/main/manager-api/src/main/java/xiaozhi/common/redis/RedisKeys.java index 0c594059..21121fd7 100644 --- a/main/manager-api/src/main/java/xiaozhi/common/redis/RedisKeys.java +++ b/main/manager-api/src/main/java/xiaozhi/common/redis/RedisKeys.java @@ -21,60 +21,9 @@ public class RedisKeys { } /** - * 登录用户Key + * 未注册设备验证码Key */ - public static String getSecurityUserKey(Long id) { - return "sys:security:user:" + id; + public static String getDeviceCaptchaKey(String captcha) { + return "sys:device:captcha:" + captcha; } - - /** - * 系统日志Key - */ - public static String getSysLogKey() { - return "sys:log"; - } - - /** - * 系统资源Key - */ - public static String getSysResourceKey() { - return "sys:resource"; - } - - /** - * 用户菜单导航Key - */ - public static String getUserMenuNavKey(Long userId) { - return "sys:user:nav:" + userId; - } - - /** - * 用户权限标识Key - */ - public static String getUserPermissionsKey(Long userId) { - return "sys:user:permissions:" + userId; - } - - /** - * 用户登陆错误次数 - * - * @param username - * @return - */ - public static String getUserLoginErrorCountKey(String username) { - return "sys:user:login:error:" + username; - } - - public static String getUserInfoKey(Long userId) { - return "sys:user:" + userId; - } - - public static String getDataScopeListKey(Long userId) { - return "sys:user:data:scope:" + userId; - } - - public static String getSysUserName(Long id) { - return "sys:user:name" + id; - } - } diff --git a/main/manager-api/src/main/java/xiaozhi/common/user/UserDetail.java b/main/manager-api/src/main/java/xiaozhi/common/user/UserDetail.java index 163982de..2ad023d7 100644 --- a/main/manager-api/src/main/java/xiaozhi/common/user/UserDetail.java +++ b/main/manager-api/src/main/java/xiaozhi/common/user/UserDetail.java @@ -3,7 +3,6 @@ package xiaozhi.common.user; import lombok.Data; import java.io.Serializable; -import java.util.List; /** * 登录用户信息 @@ -14,19 +13,7 @@ import java.util.List; public class UserDetail implements Serializable { private Long id; private String username; - private String realName; - private String headUrl; - private Integer gender; - private String email; - private String mobile; - private Long deptId; - private String password; - private Integer status; private Integer superAdmin; private String token; - /** - * 部门数据权限 - */ - private List deptIdList; - + private Integer status; } \ No newline at end of file diff --git a/main/manager-api/src/main/java/xiaozhi/common/utils/HttpContextUtils.java b/main/manager-api/src/main/java/xiaozhi/common/utils/HttpContextUtils.java index 3629d553..8d76434c 100644 --- a/main/manager-api/src/main/java/xiaozhi/common/utils/HttpContextUtils.java +++ b/main/manager-api/src/main/java/xiaozhi/common/utils/HttpContextUtils.java @@ -1,6 +1,5 @@ package xiaozhi.common.utils; -import xiaozhi.common.constant.Constant; import jakarta.servlet.http.HttpServletRequest; import org.apache.commons.lang3.StringUtils; import org.springframework.http.HttpHeaders; @@ -8,6 +7,8 @@ import org.springframework.util.DigestUtils; import org.springframework.web.context.request.RequestAttributes; import org.springframework.web.context.request.RequestContextHolder; import org.springframework.web.context.request.ServletRequestAttributes; +import xiaozhi.common.exception.ErrorCode; +import xiaozhi.common.exception.RenException; import java.util.Date; import java.util.Enumeration; @@ -30,13 +31,14 @@ public class HttpContextUtils { return ((ServletRequestAttributes) requestAttributes).getRequest(); } - public static String getToken() { - HttpServletRequest httpRequest = getHttpServletRequest(); - String token = httpRequest.getHeader(Constant.TOKEN_HEADER); - - //如果header中不存在token,则从参数中获取token + public static String getToken(String authorization) { + String token; + if (StringUtils.isBlank(authorization) && authorization.contains("Bearer ")) { + throw new RenException(ErrorCode.UNAUTHORIZED); + } + token = authorization.replace("Bearer ", ""); if (StringUtils.isBlank(token)) { - token = httpRequest.getParameter(Constant.TOKEN_HEADER); + throw new RenException(ErrorCode.TOKEN_NOT_EMPTY); } return token; } diff --git a/main/manager-api/src/main/java/xiaozhi/common/utils/Result.java b/main/manager-api/src/main/java/xiaozhi/common/utils/Result.java index c21bdbf9..65841b7d 100644 --- a/main/manager-api/src/main/java/xiaozhi/common/utils/Result.java +++ b/main/manager-api/src/main/java/xiaozhi/common/utils/Result.java @@ -1,16 +1,10 @@ package xiaozhi.common.utils; -import xiaozhi.common.exception.ErrorCode; -import xiaozhi.common.page.PageData; -import xiaozhi.modules.security.user.SecurityUser; import io.swagger.v3.oas.annotations.media.Schema; import lombok.Data; -import org.apache.commons.lang3.StringUtils; +import xiaozhi.common.exception.ErrorCode; import java.io.Serializable; -import java.util.List; -import java.util.Map; -import java.util.Set; /** * 响应数据 @@ -42,9 +36,6 @@ public class Result implements Serializable { return this; } - public boolean success() { - return code == 0; - } public Result error() { this.code = ErrorCode.INTERNAL_SERVER_ERROR; diff --git a/main/manager-api/src/main/java/xiaozhi/modules/device/controller/DeviceController.java b/main/manager-api/src/main/java/xiaozhi/modules/device/controller/DeviceController.java new file mode 100644 index 00000000..d458e9f9 --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/device/controller/DeviceController.java @@ -0,0 +1,82 @@ +package xiaozhi.modules.device.controller; + +import io.swagger.v3.oas.annotations.Operation; +import io.swagger.v3.oas.annotations.Parameter; +import io.swagger.v3.oas.annotations.Parameters; +import io.swagger.v3.oas.annotations.tags.Tag; +import lombok.AllArgsConstructor; +import org.apache.commons.lang3.StringUtils; +import org.apache.shiro.authz.annotation.RequiresPermissions; +import org.springframework.web.bind.annotation.*; +import xiaozhi.common.constant.Constant; +import xiaozhi.common.exception.ErrorCode; +import xiaozhi.common.page.PageData; +import xiaozhi.common.redis.RedisKeys; +import xiaozhi.common.redis.RedisUtils; +import xiaozhi.common.user.UserDetail; +import xiaozhi.common.utils.JsonUtils; +import xiaozhi.common.utils.Result; +import xiaozhi.modules.device.dto.DeviceHeaderDTO; +import xiaozhi.modules.device.dto.DeviceUnBindDTO; +import xiaozhi.modules.device.entity.DeviceEntity; +import xiaozhi.modules.device.service.DeviceService; +import xiaozhi.modules.security.user.SecurityUser; + +import java.util.List; +import java.util.Map; + +@Tag(name = "设备管理") +@AllArgsConstructor +@RestController +@RequestMapping("/device") +public class DeviceController { + private final DeviceService deviceService; + private final RedisUtils redisUtils; + + + @PostMapping("/bind/{deviceCode}") + @Operation(summary = "绑定设备") + @RequiresPermissions("sys:role:normal") + public Result bindDevice(@PathVariable String deviceCode) { + UserDetail user = SecurityUser.getUser(); + + String deviceHeaders = (String) redisUtils.get(RedisKeys.getDeviceCaptchaKey(deviceCode)); + if (StringUtils.isBlank(deviceHeaders)) { + return new Result().error(ErrorCode.DEVICE_CAPTCHA_ERROR); + } + DeviceHeaderDTO deviceHeader = JsonUtils.parseObject(deviceHeaders.getBytes(), DeviceHeaderDTO.class); + DeviceEntity device = deviceService.bindDevice(user.getId(), deviceHeader); + return new Result().ok(device); + } + + @GetMapping("/bind") + @Operation(summary = "获取已绑定设备") + @RequiresPermissions("sys:role:normal") + public Result> getUserDevices() { + UserDetail user = SecurityUser.getUser(); + List devices = deviceService.getUserDevices(user.getId()); + return new Result>().ok(devices); + } + + @PostMapping("/unbind") + @Operation(summary = "解绑设备") + @RequiresPermissions("sys:role:normal") + public Result unbindDevice(@RequestBody DeviceUnBindDTO unDeviveBind) { + UserDetail user = SecurityUser.getUser(); + deviceService.unbindDevice(user.getId(), unDeviveBind.getDeviceId()); + return new Result(); + } + + @GetMapping("/all") + @Operation(summary = "设备列表(管理员)") + @RequiresPermissions("sys:role:superAdmin") + @Parameters({ + @Parameter(name = Constant.PAGE, description = "当前页码,从1开始", required = true), + @Parameter(name = Constant.LIMIT, description = "每页显示记录数", required = true), + }) + public Result> adminDeviceList( + @Parameter(hidden = true) @RequestParam Map params) { + PageData page = deviceService.adminDeviceList(params); + return new Result>().ok(page); + } +} \ No newline at end of file diff --git a/main/manager-api/src/main/java/xiaozhi/modules/device/dao/DeviceDao.java b/main/manager-api/src/main/java/xiaozhi/modules/device/dao/DeviceDao.java new file mode 100644 index 00000000..9e4a8cd4 --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/device/dao/DeviceDao.java @@ -0,0 +1,9 @@ +package xiaozhi.modules.device.dao; + +import com.baomidou.mybatisplus.core.mapper.BaseMapper; +import org.apache.ibatis.annotations.Mapper; +import xiaozhi.modules.device.entity.DeviceEntity; + +@Mapper +public interface DeviceDao extends BaseMapper { +} \ No newline at end of file diff --git a/main/manager-api/src/main/java/xiaozhi/modules/device/dto/DeviceHeaderDTO.java b/main/manager-api/src/main/java/xiaozhi/modules/device/dto/DeviceHeaderDTO.java new file mode 100644 index 00000000..83327fef --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/device/dto/DeviceHeaderDTO.java @@ -0,0 +1,19 @@ +package xiaozhi.modules.device.dto; + +import io.swagger.v3.oas.annotations.media.Schema; +import lombok.Data; + +@Data +@Schema(description = "设备连接头信息") +public class DeviceHeaderDTO { + + @Schema(description = "设备ID") + private String deviceId; + + @Schema(description = "协议版本号") + private Long protocolVersion; + + @Schema(description = "认证信息") + private String authorization; + +} \ No newline at end of file diff --git a/main/manager-api/src/main/java/xiaozhi/modules/device/dto/DeviceUnBindDTO.java b/main/manager-api/src/main/java/xiaozhi/modules/device/dto/DeviceUnBindDTO.java new file mode 100644 index 00000000..830a51ae --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/device/dto/DeviceUnBindDTO.java @@ -0,0 +1,19 @@ +package xiaozhi.modules.device.dto; +import io.swagger.v3.oas.annotations.media.Schema; +import jakarta.validation.constraints.NotBlank; +import lombok.Data; + +import java.io.Serializable; + +/** + * 设备解绑表单 + */ +@Data +@Schema(description = "设备解绑表单") +public class DeviceUnBindDTO implements Serializable { + + @Schema(description = "设备ID") + @NotBlank(message = "设备ID不能为空") + private Long deviceId; + +} \ No newline at end of file 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 new file mode 100644 index 00000000..ccbe4683 --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/device/entity/DeviceEntity.java @@ -0,0 +1,54 @@ +package xiaozhi.modules.device.entity; + +import com.baomidou.mybatisplus.annotation.TableName; +import io.swagger.v3.oas.annotations.media.Schema; +import lombok.Data; + +import java.util.Date; + +@Data +@TableName("ai_device") +@Schema(description = "设备信息") +public class DeviceEntity { + @Schema(description = "设备ID") + private Long id; + + @Schema(description = "关联用户ID") + private Long userId; + + @Schema(description = "MAC地址") + private String macAddress; + + @Schema(description = "最后连接时间") + private Date lastConnectedAt; + + @Schema(description = "自动更新开关(0关闭/1开启)") + private Integer autoUpdate; + + @Schema(description = "设备硬件型号") + private String board; + + @Schema(description = "设备别名") + private String alias; + + @Schema(description = "智能体ID") + private String agentId; + + @Schema(description = "固件版本号") + private String appVersion; + + @Schema(description = "排序") + private Integer sort; + + @Schema(description = "创建者") + private Long creator; + + @Schema(description = "创建时间") + private Date createDate; + + @Schema(description = "更新者") + private Long updater; + + @Schema(description = "更新时间") + private Date updateDate; +} \ No newline at end of file diff --git a/main/manager-api/src/main/java/xiaozhi/modules/device/service/DeviceService.java b/main/manager-api/src/main/java/xiaozhi/modules/device/service/DeviceService.java new file mode 100644 index 00000000..16ad44fe --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/device/service/DeviceService.java @@ -0,0 +1,18 @@ +package xiaozhi.modules.device.service; + +import xiaozhi.common.page.PageData; +import xiaozhi.modules.device.dto.DeviceHeaderDTO; +import xiaozhi.modules.device.entity.DeviceEntity; + +import java.util.List; +import java.util.Map; + +public interface DeviceService { + DeviceEntity bindDevice(Long userId, DeviceHeaderDTO deviceHeader); + + List getUserDevices(Long userId); + + void unbindDevice(Long userId, Long deviceId); + + PageData adminDeviceList(Map params); +} \ No newline at end of file diff --git a/main/manager-api/src/main/java/xiaozhi/modules/device/service/impl/DeviceServiceImpl.java b/main/manager-api/src/main/java/xiaozhi/modules/device/service/impl/DeviceServiceImpl.java new file mode 100644 index 00000000..279a9f2e --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/device/service/impl/DeviceServiceImpl.java @@ -0,0 +1,57 @@ +package xiaozhi.modules.device.service.impl; + +import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper; +import com.baomidou.mybatisplus.core.metadata.IPage; +import org.springframework.stereotype.Service; +import xiaozhi.common.page.PageData; +import xiaozhi.common.service.impl.BaseServiceImpl; +import xiaozhi.modules.device.dao.DeviceDao; +import xiaozhi.modules.device.dto.DeviceHeaderDTO; +import xiaozhi.modules.device.entity.DeviceEntity; +import xiaozhi.modules.device.service.DeviceService; + +import java.util.Date; +import java.util.List; +import java.util.Map; + +@Service +public class DeviceServiceImpl extends BaseServiceImpl implements DeviceService { + private final DeviceDao deviceDao; + + // 添加构造函数来初始化 deviceMapper + public DeviceServiceImpl(DeviceDao deviceDao) { + this.deviceDao = deviceDao; + } + + @Override + public DeviceEntity bindDevice(Long userId, DeviceHeaderDTO deviceHeader) { + DeviceEntity device = new DeviceEntity(); + device.setUserId(userId); + device.setMacAddress(deviceHeader.getDeviceId()); + device.setCreateDate(new Date()); + deviceDao.insert(device); + return device; + } + + @Override + public List getUserDevices(Long userId) { + QueryWrapper wrapper = new QueryWrapper<>(); + wrapper.eq("user_id", userId); + return deviceDao.selectList(wrapper); + } + + @Override + public void unbindDevice(Long userId, Long deviceId) { + deviceDao.deleteById(deviceId); + } + + @Override + public PageData adminDeviceList(Map params) { + IPage page = deviceDao.selectPage( + getPage(params, "sort", true), + new QueryWrapper<>() + ); + return new PageData<>(page.getRecords(), page.getTotal()); + } + +} \ No newline at end of file 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 29f6b706..de25d80e 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,18 +58,21 @@ 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"); - filterMap.put("/login", "anon"); - filterMap.put("/publicKey", "anon"); filterMap.put("/v3/api-docs/**", "anon"); filterMap.put("/doc.html", "anon"); - filterMap.put("/sys/oss/download/**", "anon"); - filterMap.put("/captcha", "anon"); filterMap.put("/favicon.ico", "anon"); - filterMap.put("/mobile/**", "anon"); - filterMap.put("/user/**", "anon"); + filterMap.put("/user/captcha", "anon"); + filterMap.put("/user/login", "anon"); + filterMap.put("/user/register", "anon"); filterMap.put("/**", "oauth2"); shiroFilter.setFilterChainDefinitionMap(filterMap); 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..6ab9d1b9 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,15 +4,21 @@ 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.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.user.UserDetail; import xiaozhi.common.utils.Result; import xiaozhi.common.validator.AssertUtils; import xiaozhi.modules.security.dto.LoginDTO; import xiaozhi.modules.security.password.PasswordUtils; import xiaozhi.modules.security.service.CaptchaService; import xiaozhi.modules.security.service.SysUserTokenService; +import xiaozhi.modules.security.user.SecurityUser; +import xiaozhi.modules.sys.dto.PasswordDTO; import xiaozhi.modules.sys.dto.SysUserDTO; import xiaozhi.modules.sys.service.SysUserService; @@ -21,7 +27,6 @@ import java.io.IOException; /** * 登录控制层 */ -@Tag(name = "登录管理") @AllArgsConstructor @RestController @RequestMapping("/user") @@ -44,7 +49,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) { @@ -73,15 +78,31 @@ public class LoginController { } // 按照用户名获取用户 SysUserDTO userDTO = sysUserService.getByUsername(login.getUsername()); - if (userDTO != null){ + if (userDTO != null) { throw new RenException("此手机号码已经注册过"); } userDTO = new SysUserDTO(); userDTO.setUsername(login.getUsername()); userDTO.setPassword(login.getPassword()); sysUserService.save(userDTO); - return new Result(); + return new Result<>(); } + @GetMapping("/info") + @Operation(summary = "用户信息获取") + public Result info() { + UserDetail user = SecurityUser.getUser(); + Result result = new Result<>(); + result.setData(user); + return result; + } + + @PutMapping("/change-password") + @Operation(summary = "修改用户密码") + public Result changePassword(@RequestBody PasswordDTO passwordDTO) { + Long userId = SecurityUser.getUserId(); + sysUserTokenService.changePassword(userId, passwordDTO); + return new Result<>(); + } } \ 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..1520a850 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 @@ -89,12 +89,11 @@ public class Oauth2Filter extends AuthenticatingFilter { * 获取请求的token */ private String getRequestToken(HttpServletRequest httpRequest) { + String token = null; //从header中获取token - String token = httpRequest.getHeader(Constant.TOKEN_HEADER); - - //如果header中不存在token,则从参数中获取token - if (StringUtils.isBlank(token)) { - token = httpRequest.getParameter(Constant.TOKEN_HEADER); + String authorization = httpRequest.getHeader(Constant.AUTHORIZATION); + if (StringUtils.isNotBlank(authorization) && authorization.startsWith("Bearer ")) { + 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 933724a1..b7d6c278 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 @@ -1,5 +1,15 @@ package xiaozhi.modules.security.oauth2; +import jakarta.annotation.Resource; +import org.apache.shiro.authc.*; +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; @@ -7,16 +17,8 @@ import xiaozhi.common.utils.MessageUtils; import xiaozhi.modules.security.entity.SysUserTokenEntity; import xiaozhi.modules.security.service.ShiroService; import xiaozhi.modules.sys.entity.SysUserEntity; -import jakarta.annotation.Resource; -import org.apache.shiro.authc.*; -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.springframework.context.annotation.Lazy; -import org.springframework.stereotype.Component; +import xiaozhi.modules.sys.enums.SuperAdminEnum; -import java.util.List; import java.util.Set; /** @@ -30,6 +32,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; @@ -45,6 +49,13 @@ public class Oauth2Realm extends AuthorizingRealm { //用户权限列表 Set permsSet = shiroService.getUserPermissions(user); + if (user.getSuperAdmin() == SuperAdminEnum.YES.value()) { + permsSet.add("sys:role:superAdmin"); + permsSet.add("sys:role:normal"); + } else { + permsSet.add("sys:role:normal"); + } + SimpleAuthorizationInfo info = new SimpleAuthorizationInfo(); info.setStringPermissions(permsSet); return info; @@ -70,11 +81,14 @@ public class Oauth2Realm extends AuthorizingRealm { //转换成UserDetail对象 UserDetail userDetail = ConvertUtils.sourceToTarget(userEntity, UserDetail.class); - //获取用户对应的部门数据权限 - userDetail.setDeptIdList(null); 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..d1d094d4 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,11 +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 java.util.Map; +import xiaozhi.modules.sys.dto.PasswordDTO; +import xiaozhi.modules.sys.dto.SysUserDTO; /** * 用户Token @@ -19,7 +19,9 @@ public interface SysUserTokenService extends BaseService { * * @param userId 用户ID */ - Result createToken(Long userId); + Result createToken(Long userId); + + SysUserDTO getUserByToken(String token); /** * 退出 @@ -28,4 +30,12 @@ public interface SysUserTokenService extends BaseService { */ void logout(Long userId); + /** + * 修改密码 + * + * @param userId + * @param passwordDTO + */ + void changePassword(Long userId, PasswordDTO passwordDTO); + } \ No newline at end of file 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..40ee7957 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,10 @@ package xiaozhi.modules.security.service.impl; import cn.hutool.core.date.DateUtil; +import lombok.AllArgsConstructor; +import org.springframework.stereotype.Service; +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; @@ -9,19 +13,24 @@ import xiaozhi.modules.security.dao.SysUserTokenDao; 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.PasswordDTO; +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,9 +79,36 @@ public class SysUserTokenServiceImpl extends BaseServiceImpl> page(@Parameter(hidden = true) @RequestParam Map params) { - PageData page = sysUserService.page(params); - - return new Result>().ok(page); - } - - @GetMapping("{id}") - @Operation(summary = "信息") - @RequiresPermissions("sys:user:info") - public Result get(@PathVariable("id") Long id) { - SysUserDTO data = sysUserService.get(id); - return new Result().ok(data); - } - - @GetMapping("info") - @Operation(summary = "登录用户信息") - public Result info() { - SysUserDTO data = ConvertUtils.sourceToTarget(SecurityUser.getUser(), SysUserDTO.class); - return new Result().ok(data); - } - - @PutMapping("password") - @Operation(summary = "修改密码") - @LogOperation("修改密码") - public Result password(@RequestBody PasswordDTO dto) { - //效验数据 - ValidatorUtils.validateEntity(dto); - String newPassword = dto.getNewPassword(); - - //密码的强度 - if (newPassword == null || newPassword.length() < 8) { - return new Result().error(ErrorCode.PASSWORD_LENGTH_ERROR); - } - if (!sysUserService.isStrongPassword(newPassword)) { - return new Result().error(ErrorCode.PASSWORD_WEAK_ERROR); - } - UserDetail user = SecurityUser.getUser(); - //原密码不正确 - if (!PasswordUtils.matches(dto.getPassword(), user.getPassword())) { - return new Result().error(ErrorCode.PASSWORD_ERROR); - } - - sysUserService.updatePassword(user.getId(), dto.getNewPassword()); - - return new Result(); - } - - @PostMapping - @Operation(summary = "保存") - @LogOperation("保存") - @RequiresPermissions("sys:user:save") - public Result save(@RequestBody SysUserDTO dto) { - //效验数据 - ValidatorUtils.validateEntity(dto, AddGroup.class, DefaultGroup.class); - - sysUserService.save(dto); - - return new Result(); - } - - @PutMapping - @Operation(summary = "修改") - @LogOperation("修改") - @RequiresPermissions("sys:user:update") - public Result update(@RequestBody SysUserDTO dto) { - //效验数据 - ValidatorUtils.validateEntity(dto, UpdateGroup.class, DefaultGroup.class); - - sysUserService.update(dto); - - return new Result(); - } - - @PutMapping("app") - @Operation(summary = "修改") - @LogOperation("修改") - @RequiresPermissions("sys:user:update") - public Result updateUserInfo(@RequestBody SysUserDTO dto) { - sysUserService.updateUserInfo(dto); - - return new Result(); - } - - @DeleteMapping - @Operation(summary = "删除") - @LogOperation("删除") - @RequiresPermissions("sys:user:delete") - public Result delete(@RequestBody Long[] ids) { - //效验数据 - AssertUtils.isArrayEmpty(ids, "id"); - - List idList = Arrays.asList(ids); - if (idList.contains(SecurityUser.getUserId())) { - throw new RenException(ErrorCode.DEL_MYSELF_ERROR); - } - - sysUserService.deleteBatchIds(idList); - - return new Result(); - } } \ No newline at end of file diff --git a/main/manager-api/src/main/java/xiaozhi/modules/sys/dao/SysUserDao.java b/main/manager-api/src/main/java/xiaozhi/modules/sys/dao/SysUserDao.java index 08c48a4d..339e522a 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/sys/dao/SysUserDao.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/sys/dao/SysUserDao.java @@ -1,12 +1,8 @@ package xiaozhi.modules.sys.dao; +import org.apache.ibatis.annotations.Mapper; import xiaozhi.common.dao.BaseDao; import xiaozhi.modules.sys.entity.SysUserEntity; -import org.apache.ibatis.annotations.Mapper; -import org.apache.ibatis.annotations.Param; - -import java.util.List; -import java.util.Map; /** * 系统用户 @@ -14,21 +10,4 @@ import java.util.Map; @Mapper public interface SysUserDao extends BaseDao { - List getList(Map params); - - SysUserEntity getById(Long id); - - SysUserEntity getByUsername(String username); - - int updatePassword(@Param("id") Long id, @Param("newPassword") String newPassword); - - /** - * 根据部门ID,查询用户数 - */ - int getCountByDeptId(Long deptId); - - /** - * 根据部门ID,查询用户ID列表 - */ - List getUserIdListByDeptId(List deptIdList); } \ No newline at end of file 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 a5d14a71..0093fcf0 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 @@ -3,9 +3,9 @@ package xiaozhi.modules.sys.entity; import com.baomidou.mybatisplus.annotation.FieldFill; import com.baomidou.mybatisplus.annotation.TableField; import com.baomidou.mybatisplus.annotation.TableName; -import xiaozhi.common.entity.BaseEntity; import lombok.Data; import lombok.EqualsAndHashCode; +import xiaozhi.common.entity.BaseEntity; import java.util.Date; @@ -24,30 +24,6 @@ public class SysUserEntity extends BaseEntity { * 密码 */ private String password; - /** - * 姓名 - */ - private String realName; - /** - * 头像 - */ - private String headUrl; - /** - * 性别 0:男 1:女 2:保密 - */ - private Integer gender; - /** - * 邮箱 - */ - private String email; - /** - * 手机号 - */ - private String mobile; - /** - * 部门ID - */ - private Long deptId; /** * 超级管理员 0:否 1:是 */ @@ -66,10 +42,5 @@ public class SysUserEntity extends BaseEntity { */ @TableField(fill = FieldFill.INSERT_UPDATE) private Date updateDate; - /** - * 部门名称 - */ - @TableField(exist = false) - private String deptName; } \ No newline at end of file diff --git a/main/manager-api/src/main/java/xiaozhi/modules/sys/service/SysUserService.java b/main/manager-api/src/main/java/xiaozhi/modules/sys/service/SysUserService.java index 8427fd09..08179409 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/sys/service/SysUserService.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/sys/service/SysUserService.java @@ -1,49 +1,23 @@ package xiaozhi.modules.sys.service; -import xiaozhi.common.page.PageData; import xiaozhi.common.service.BaseService; +import xiaozhi.modules.sys.dto.PasswordDTO; import xiaozhi.modules.sys.dto.SysUserDTO; import xiaozhi.modules.sys.entity.SysUserEntity; -import java.util.List; -import java.util.Map; - /** * 系统用户 */ public interface SysUserService extends BaseService { - PageData page(Map params); - - List list(Map params); - - SysUserDTO get(Long id); - SysUserDTO getByUsername(String username); + SysUserDTO getByUserId(Long userId); + void save(SysUserDTO dto); - void update(SysUserDTO dto); - - void updateUserInfo(SysUserDTO dto); - void delete(Long[] ids); - /** - * 修改密码 - * - * @param id 用户ID - * @param newPassword 新密码 - */ - void updatePassword(Long id, String newPassword); - - /** - * 验证密码强度 - * - * @param newPassword - * @return - */ - boolean isStrongPassword(String newPassword); - + void changePassword(Long userId, PasswordDTO passwordDTO); } 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 1828039b..17d1939c 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 @@ -1,19 +1,16 @@ package xiaozhi.modules.sys.service.impl; import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper; -import com.baomidou.mybatisplus.core.metadata.IPage; import lombok.AllArgsConstructor; -import org.apache.commons.lang3.StringUtils; import org.springframework.stereotype.Service; import org.springframework.transaction.annotation.Transactional; import xiaozhi.common.exception.ErrorCode; import xiaozhi.common.exception.RenException; -import xiaozhi.common.page.PageData; -import xiaozhi.common.redis.RedisUtils; import xiaozhi.common.service.impl.BaseServiceImpl; import xiaozhi.common.utils.ConvertUtils; import xiaozhi.modules.security.password.PasswordUtils; import xiaozhi.modules.sys.dao.SysUserDao; +import xiaozhi.modules.sys.dto.PasswordDTO; import xiaozhi.modules.sys.dto.SysUserDTO; import xiaozhi.modules.sys.entity.SysUserEntity; import xiaozhi.modules.sys.enums.SuperAdminEnum; @@ -21,7 +18,6 @@ import xiaozhi.modules.sys.service.SysUserService; import java.util.Arrays; import java.util.List; -import java.util.Map; import java.util.regex.Matcher; import java.util.regex.Pattern; @@ -32,43 +28,27 @@ import java.util.regex.Pattern; @AllArgsConstructor @Service public class SysUserServiceImpl extends BaseServiceImpl implements SysUserService { - private final RedisUtils redisUtils; - - @Override - public PageData page(Map params) { - //转换成like - paramsToLike(params, "username"); - - //分页 - IPage page = getPage(params, "t1.create_date", false); - - //查询 - List list = baseDao.getList(params); - - return getPageData(list, page.getTotal(), SysUserDTO.class); - } - - @Override - public List list(Map params) { - - List entityList = baseDao.getList(params); - - return ConvertUtils.sourceToTarget(entityList, SysUserDTO.class); - } - - @Override - public SysUserDTO get(Long id) { - SysUserEntity entity = baseDao.getById(id); - - return ConvertUtils.sourceToTarget(entity, SysUserDTO.class); - } + private final SysUserDao sysUserDao; @Override public SysUserDTO getByUsername(String username) { - SysUserEntity entity = baseDao.getByUsername(username); + QueryWrapper queryWrapper = new QueryWrapper<>(); + queryWrapper.eq("username", username); + List users = sysUserDao.selectList(queryWrapper); + if (users == null || users.isEmpty()) { + return null; + } + SysUserEntity entity = users.getFirst(); return ConvertUtils.sourceToTarget(entity, SysUserDTO.class); } + @Override + public SysUserDTO getByUserId(Long userId) { + SysUserEntity sysUserEntity = sysUserDao.selectById(userId); + + return ConvertUtils.sourceToTarget(sysUserEntity, SysUserDTO.class); + } + @Override @Transactional(rollbackFor = Exception.class) public void save(SysUserDTO dto) { @@ -90,44 +70,11 @@ public class SysUserServiceImpl extends BaseServiceImpl - - - - - - - - update sys_user set password = #{newPassword} where id = #{id} - - \ No newline at end of file diff --git a/main/manager-web/src/App.vue b/main/manager-web/src/App.vue index c4f1de8c..9c7e6db7 100644 --- a/main/manager-web/src/App.vue +++ b/main/manager-web/src/App.vue @@ -26,4 +26,7 @@ nav { } } } + + \ 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..e7ba492e --- /dev/null +++ b/main/manager-web/src/components/AddDeviceDialog.vue @@ -0,0 +1,90 @@ + + + + + \ No newline at end of file diff --git a/main/manager-web/src/components/AddWisdomBodyDialog.vue b/main/manager-web/src/components/AddWisdomBodyDialog.vue new file mode 100644 index 00000000..fef13d50 --- /dev/null +++ b/main/manager-web/src/components/AddWisdomBodyDialog.vue @@ -0,0 +1,87 @@ + + + + + \ 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..e830555f --- /dev/null +++ b/main/manager-web/src/components/DeviceItem.vue @@ -0,0 +1,87 @@ + + + + \ 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..758959b3 --- /dev/null +++ b/main/manager-web/src/components/HeaderBar.vue @@ -0,0 +1,135 @@ + + + + + \ 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..d3add00a 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,35 +7,47 @@ 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') } }, + // 新增设备管理页面路由 + { + path: '/device-management', + name: 'DeviceManagement', + component: function () { + return import('../views/DeviceManagement.vue') + } + + } ] const router = new VueRouter({ 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/DeviceManagement.vue b/main/manager-web/src/views/DeviceManagement.vue new file mode 100644 index 00000000..06c4317c --- /dev/null +++ b/main/manager-web/src/views/DeviceManagement.vue @@ -0,0 +1,169 @@ + + + + + \ No newline at end of file diff --git a/main/manager-web/src/views/home.vue b/main/manager-web/src/views/home.vue index 4d0181f5..4b9c6eb5 100644 --- a/main/manager-web/src/views/home.vue +++ b/main/manager-web/src/views/home.vue @@ -1,347 +1,94 @@ - - - + \ 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..77ecc85f --- /dev/null +++ b/main/manager-web/src/views/roleConfig.vue @@ -0,0 +1,294 @@ + + + + + + diff --git a/main/xiaozhi-server/config.yaml b/main/xiaozhi-server/config.yaml index a2fa3446..f3ca5706 100644 --- a/main/xiaozhi-server/config.yaml +++ b/main/xiaozhi-server/config.yaml @@ -35,10 +35,7 @@ log: log_file: "server.log" # 设置数据文件路径 data_dir: data -iot: - Speaker: - # 设置esp32的音量,范围0-100 - volume: 80 + xiaozhi: type: hello version: 1 @@ -51,13 +48,15 @@ xiaozhi: prompt: | 你是一个叫小智/小志的台湾女孩,说话机车,声音好听,习惯简短表达,爱用网络梗。 请注意,要像一个人一样说话,请不要回复表情符号、代码、和xml标签。 - 当前时间是:{date_time},现在我正在和你进行语音聊天,我们开始吧。 + 现在我正在和你进行语音聊天,我们开始吧。 如果用户希望结束对话,请在最后说“拜拜”或“再见”。 # 使用完声音文件后删除文件(Delete the sound file when you are done using it) delete_audio: true # 没有语音输入多久后断开连接(秒),默认2分钟,即120秒 close_connection_no_voice_time: 120 +# TTS请求超时时间(秒) +tts_timeout: 10 CMD_exit: - "退出" @@ -85,14 +84,28 @@ selected_module: Intent: # 不使用意图识别 nointent: - # 不需要动 + # 不需要动type type: nointent intent_llm: - # 不需要动 + # 不需要动type type: intent_llm function_call: - # 不需要动 + # 不需要动type type: nointent + # plugins_func/functions下的模块,可以通过配置,选择加载哪个模块,加载后对话支持相应的function调用 + # 系统默认已经记载“handle_exit_intent(退出识别)”、“play_music(音乐播放)”插件,请勿重复加载 + # 下面是加载查天气、角色切换的插件示例 + functions: + - change_role + - get_weather + +# 插件的基础配置 +plugins: + # 获取天气插件的配置,这里填写你的api_key + # 这个密钥是项目共用的key,用多了可能会被限制 + # 想稳定一点就自行申请替换,每天有1000次免费调用 + # 申请地址:https://console.qweather.com/#/apps/create-key/over + get_weather: { "api_key": "a861d0d5e7bf4ee1a83d9a9e4f96d4da", "default_location": "广州" } Memory: mem0ai: @@ -206,18 +219,13 @@ LLM: variables: k: "v" k2: "v2" - -TTS_SET: - TTS_STREAM: false #是否启动流响应 true/false,默认false - MAX_WORKERS: 4 #并发请求tts的数量,这个调太大,本地tts会变慢 - TTS: # 当前支持的type为edge、doubao,可自行适配 EdgeTTS: # 定义TTS API类型 type: edge voice: zh-CN-XiaoxiaoNeural - output_file: tmp/ + output_dir: tmp/ DoubaoTTS: # 定义TTS API类型 type: doubao @@ -227,7 +235,7 @@ TTS: # 地址:https://console.volcengine.com/speech/service/8 api_url: https://openspeech.bytedance.com/api/v1/tts voice: BV001_streaming - output_file: tmp/ + output_dir: tmp/ authorization: "Bearer;" appid: 你的火山引擎语音合成服务appid access_token: 你的火山引擎语音合成服务access_token @@ -238,7 +246,7 @@ TTS: # token申请地址 https://cloud.siliconflow.cn/account/ak model: FunAudioLLM/CosyVoice2-0.5B voice: FunAudioLLM/CosyVoice2-0.5B:alex - output_file: tmp/ + output_dir: tmp/ access_token: 你的硅基流动API密钥 response_format: wav CozeCnTTS: @@ -246,7 +254,7 @@ TTS: # COZECN TTS # token申请地址 https://www.coze.cn/open/oauth/pats voice: 7426720361733046281 - output_file: tmp/ + output_dir: tmp/ access_token: 你的coze web key response_format: wav FishSpeech: @@ -259,26 +267,18 @@ TTS: #--decoder-config-name firefly_gan_vq #--compile type: fishspeech - output_file: tmp/ + output_dir: tmp/ response_format: wav reference_id: null - reference_audio: - - "audio_ref/shinchan.wav" - - "audio_ref/shinchan2.wav" - - "audio_ref/shinchan3.wav" - - "audio_ref/shinchan4.wav" - reference_text: - - "难过的事讨厌的事丢脸的事全都集中起来,在这里,让水冲走所有的烦恼不就好了吗。" - - "我还是觉得船到桥头自然直这句话最棒了。" - - "我不是小鬼,我是野原新之助,小新这次的目标是成为圆梦之星哦,在努力一下好了,是我的脚自己要走这么快的喔,竞争就是这么激烈。小白,我们走了,成为圆梦之星是不是可以打败怪兽呢。" - - "工钱,今天一起算,有磨能使鬼推钱。哎呀,你难倒我了,像我这么乖的小孩,怎么更乖。" + reference_audio: ["/tmp/test.wav",] + reference_text: ["你弄来这些吟词宴曲来看,还是这些混话来欺负我。",] normalize: true max_new_tokens: 1024 chunk_length: 200 top_p: 0.7 repetition_penalty: 1.2 temperature: 0.7 - streaming: true + streaming: false use_memory_cache: "on" seed: null channels: 1 @@ -291,7 +291,7 @@ TTS: #python api_v2.py -a 127.0.0.1 -p 9880 -c GPT_SoVITS/configs/caixukun.yaml type: gpt_sovits_v2 url: "http://127.0.0.1:9880/tts" - output_file: tmp/ + output_dir: tmp/ text_lang: "auto" ref_audio_path: "caixukun.wav" prompt_text: "" @@ -311,20 +311,24 @@ TTS: repetition_penalty: 1.35 aux_ref_audio_paths: [] GPT_SOVITS_V3: + # 定义TTS API类型 GPT-SoVITS-v3lora-20250228 + #启动tts方法: + #python api.py type: gpt_sovits_v3 - url: "http://127.0.0.1:9880/tts" - output_file: tmp/ - text_lang: "auto" - ref_audio_path: "caixukun.wav" - prompt_lang: "zh" + url: "http://127.0.0.1:9880" + output_dir: tmp/ + text_language: "auto" + refer_wav_path: "caixukun.wav" + prompt_language: "zh" prompt_text: "" - top_k: 5 - top_p: 1 - temperature: 1 - sample_steps: 16 - media_type: "wav" - streaming_mode: false - threshold: 30 + top_k: 15 + top_p: 1.0 + temperature: 1.0 + cut_punc: "" + speed: 1.0 + inp_refs: [] + sample_steps: 32 + if_sr: false MinimaxTTS: # Minimax语音合成服务,需要先在minimax平台创建账户充值,并获取登录信息 # 平台地址:https://platform.minimaxi.com/ @@ -333,7 +337,7 @@ TTS: # api_key地址:https://platform.minimaxi.com/user-center/basic-information/interface-key # 定义TTS API类型 type: minimax - output_file: tmp/ + output_dir: tmp/ group_id: 你的minimax平台groupID api_key: 你的minimax平台接口密钥 model: "speech-01-turbo" @@ -370,7 +374,7 @@ TTS: # token地址:https://nls-portal.console.aliyun.com/overview # 定义TTS API类型 type: aliyun - output_file: tmp/ + output_dir: tmp/ appkey: 你的阿里云智能语音交互服务项目Appkey token: 你的阿里云智能语音交互服务AccessToken,临时的24小时,要长期用下方的access_key_id,access_key_secret voice: xiaoyun @@ -393,7 +397,7 @@ TTS: api_url: https://api.302ai.cn/doubao/tts_hd authorization: "Bearer " voice: "zh_female_wanwanxiaohe_moon_bigtts" - output_file: tmp/ + output_dir: tmp/ access_token: "你的302API密钥" ACGNTTS: #在线网址:https://acgn.ttson.cn/ @@ -410,7 +414,7 @@ TTS: to_lang: ZH url: https://u95167-bd74-2aef8085.westx.seetacloud.com:8443/flashsummary/tts?token= format: mp3 - output_file: tmp/ + output_dir: tmp/ emotion: 1 OpenAITTS: # openai官方文本转语音服务,可支持全球大多数语种 @@ -424,7 +428,22 @@ TTS: voice: onyx # 语速范围0.25-4.0 speed: 1 - output_file: tmp/ + output_dir: 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_dir: tmp/ # 模块测试配置 module_test: test_sentences: # 自定义测试语句 @@ -447,4 +466,4 @@ manager: enabled: false ip: 0.0.0.0 port: 8002 -use_private_config: false \ No newline at end of file +use_private_config: false diff --git a/main/xiaozhi-server/config/functionCallConfig.py b/main/xiaozhi-server/config/functionCallConfig.py deleted file mode 100644 index dd9b55fc..00000000 --- a/main/xiaozhi-server/config/functionCallConfig.py +++ /dev/null @@ -1,36 +0,0 @@ -FunctionCallConfig = [ - { - "type": "function", - "function": { - "name": "handle_exit_intent", - "description": "当用户想结束对话或需要退出系统时调用", - "parameters": { - "type": "object", - "properties": { - "say_goodbye": { - "type": "string", - "description": "和用户友好结束对话的告别语" - } - }, - "required": [] - } - } - }, - { - "type": "function", - "function": { - "name": "play_music", - "description": "唱歌、听歌、播放音乐方法。比如用户说播放音乐,参数为:random,比如用户说播放两只老虎,参数为:两只老虎", - "parameters": { - "type": "object", - "properties": { - "song_name": { - "type": "string", - "description": "歌曲名称,如果没有指定具体歌名则为'random'" - } - }, - "required": ["song_name"] - } - } - } - ] \ No newline at end of file diff --git a/main/xiaozhi-server/config/settings.py b/main/xiaozhi-server/config/settings.py index 0582bf76..bb81cfe9 100644 --- a/main/xiaozhi-server/config/settings.py +++ b/main/xiaozhi-server/config/settings.py @@ -7,9 +7,48 @@ from core.utils.util import read_config, get_project_dir default_config_file = "config.yaml" +def ensure_directories(config): + """确保所有配置路径存在""" + dirs_to_create = set() + project_dir = get_project_dir() # 获取项目根目录 + # 日志文件目录 + log_dir = config.get('log', {}).get('log_dir', 'tmp') + dirs_to_create.add(os.path.join(project_dir, log_dir)) + + # ASR/TTS模块输出目录 + for module in ['ASR', 'TTS']: + for provider in config.get(module, {}).values(): + output_dir = provider.get('output_dir', '') + if output_dir: + dirs_to_create.add(output_dir) + + # 根据selected_module创建模型目录 + selected_modules = config.get('selected_module', {}) + for module_type in ['ASR', 'LLM', 'TTS']: + selected_provider = selected_modules.get(module_type) + if not selected_provider: + continue + provider_config = config.get(module_type, {}).get(selected_provider, {}) + output_dir = provider_config.get('output_dir') + if output_dir: + full_model_dir = os.path.join(project_dir, output_dir) + dirs_to_create.add(full_model_dir) + + # 统一创建目录(保留原data目录创建) + for dir_path in dirs_to_create: + try: + os.makedirs(dir_path, exist_ok=True) + except PermissionError: + print(f"警告:无法创建目录 {dir_path},请检查写入权限") + + def get_config_file(): global default_config_file - # 判断是否存在私有的配置文件 + """获取配置文件路径,优先使用私有配置文件(若存在)。 + + Returns: + str: 配置文件路径(相对路径或默认路径) + """ config_file = default_config_file if os.path.exists(get_project_dir() + "data/." + default_config_file): config_file = "data/." + default_config_file @@ -20,9 +59,13 @@ def load_config(): """加载配置文件""" parser = argparse.ArgumentParser(description="Server configuration") config_file = get_config_file() + parser.add_argument("--config_path", type=str, default=config_file) args = parser.parse_args() - return read_config(args.config_path) + config = read_config(args.config_path) + # 初始化目录 + ensure_directories(config) + return config def update_config(config): @@ -67,7 +110,7 @@ def find_missing_keys(new_config, old_config, parent_key=''): def check_config_file(): old_config_file = get_config_file() global default_config_file - if not old_config_file.startswith('data'): + if not 'data' in old_config_file: return old_config = read_config(get_project_dir() + old_config_file) new_config = read_config(get_project_dir() + default_config_file) diff --git a/main/xiaozhi-server/core/connection.py b/main/xiaozhi-server/core/connection.py index f53aa1e8..1e4d0fac 100644 --- a/main/xiaozhi-server/core/connection.py +++ b/main/xiaozhi-server/core/connection.py @@ -5,17 +5,20 @@ import time import queue import asyncio import traceback -from config.logger import setup_logging + import threading import websockets from typing import Dict, Any +import plugins_func.loadplugins +from config.logger import setup_logging from core.utils.dialogue import Message, Dialogue from core.handle.textHandle import handleTextMessage -from core.utils.util import get_string_no_punctuation_or_emoji, extract_json_from_string +from core.utils.util import get_string_no_punctuation_or_emoji, extract_json_from_string, get_ip_info from concurrent.futures import ThreadPoolExecutor, TimeoutError from core.handle.sendAudioHandle import sendAudioMessage, sendAudioMessageStream from core.handle.receiveAudioHandle import handleAudioMessage -from core.handle.intentHandler import Action, get_functions, handle_llm_function_call +from core.handle.functionHandler import FunctionHandler +from plugins_func.register import Action from config.private_config import PrivateConfig from core.auth import AuthMiddleware, AuthenticationError from core.utils.auth_code_gen import AuthCodeGenerator @@ -28,7 +31,7 @@ class TTSException(RuntimeError): class ConnectionHandler: - def __init__(self, config: Dict[str, Any], _vad, _asr, _llm, _tts, _music, _memory, _intent): + def __init__(self, config: Dict[str, Any], _vad, _asr, _llm, _tts, _memory, _intent): self.config = config self.logger = setup_logging() self.auth = AuthMiddleware(config) @@ -37,6 +40,8 @@ class ConnectionHandler: self.websocket = None self.headers = None + self.client_ip = None + self.client_ip_info = {} self.session_id = None self.prompt = None self.welcome_msg = None @@ -94,7 +99,6 @@ class ConnectionHandler: self.private_config = None self.auth_code_gen = AuthCodeGenerator.get_instance() self.is_device_verified = False # 添加设备验证状态标志 - self.music_handler = _music self.close_after_chat = False # 是否在聊天结束后关闭连接 self.use_function_call_mode = False if self.config["selected_module"]["Intent"] == 'function_call': @@ -105,8 +109,8 @@ class ConnectionHandler: # 获取并验证headers self.headers = dict(ws.request.headers) # 获取客户端ip地址 - client_ip = ws.remote_address[0] - self.logger.bind(tag=TAG).info(f"{client_ip} conn - Headers: {self.headers}") + self.client_ip = ws.remote_address[0] + self.logger.bind(tag=TAG).info(f"{self.client_ip} conn - Headers: {self.headers}") # 进行认证 await self.auth.authenticate(self.headers) @@ -150,6 +154,7 @@ class ConnectionHandler: self.welcome_msg["session_id"] = self.session_id await self.websocket.send(json.dumps(self.welcome_msg)) + # 异步初始化 await self.loop.run_in_executor(None, self._initialize_components) # tts 消化线程 @@ -190,12 +195,21 @@ class ConnectionHandler: self.prompt = self.config["prompt"] if self.private_config: self.prompt = self.private_config.private_config.get("prompt", self.prompt) - # 赋予LLM时间观念 - if "{date_time}" in self.prompt: - date_time = time.strftime("%Y-%m-%d %H:%M", time.localtime()) - self.prompt = self.prompt.replace("{date_time}", date_time) + + self.client_ip_info = get_ip_info(self.client_ip) + self.logger.bind(tag=TAG).info(f"Client ip info: {self.client_ip_info}") + self.prompt = self.prompt + f"\n我在:{self.client_ip_info}" self.dialogue.put(Message(role="system", content=self.prompt)) + self.func_handler = FunctionHandler(self.config) + + def change_system_prompt(self, prompt): + self.prompt = prompt + # 找到原来的role==system,替换原来的系统提示 + for m in self.dialogue.dialogue: + if m.role == "system": + m.content = prompt + async def _check_and_broadcast_auth_code(self): """检查设备绑定状态并广播认证码""" if not self.private_config.get_owner(): @@ -312,7 +326,7 @@ class ConnectionHandler: self.logger.bind(tag=TAG).debug(json.dumps(self.dialogue.get_llm_dialogue(), indent=4, ensure_ascii=False)) return True - def chat_with_function_calling(self, query): + def chat_with_function_calling(self, query, tool_call=False): self.logger.bind(tag=TAG).debug(f"Chat with function calling start: {query}") """Chat with function calling for intent detection using streaming""" if self.isNeedAuth(): @@ -321,10 +335,11 @@ class ConnectionHandler: future.result() return True - self.dialogue.put(Message(role="user", content=query)) + if not tool_call: + self.dialogue.put(Message(role="user", content=query)) # Define intent functions - functions = get_functions() + functions = self.func_handler.get_functions() response_message = [] processed_chars = 0 # 跟踪已处理的字符位置 @@ -360,7 +375,7 @@ class ConnectionHandler: for response in llm_responses: content, tools_call = response if content is not None and len(content) > 0: - if len(response_message) <= 0 and content == "```": + if len(response_message) <= 0 and (content == "```" or "" in content): tool_call_flag = True if tools_call is not None: @@ -417,6 +432,38 @@ class ConnectionHandler: self.tts_queue.put(future) processed_chars += len(segment_text_raw) # 更新已处理字符位置 + # 处理function call + if tool_call_flag: + bHasError = False + if function_id is None: + a = extract_json_from_string(content_arguments) + if a is not None: + try: + content_arguments_json = json.loads(a) + function_name = content_arguments_json["name"] + function_arguments = json.dumps(content_arguments_json["arguments"], ensure_ascii=False) + function_id = str(uuid.uuid4().hex) + except Exception as e: + bHasError = True + response_message.append(a) + else: + bHasError = True + response_message.append(content_arguments) + if bHasError: + self.logger.bind(tag=TAG).error(f"function call error: {content_arguments}") + else: + function_arguments = json.loads(function_arguments) + if not bHasError: + self.logger.bind(tag=TAG).info( + f"function_name={function_name}, function_id={function_id}, function_arguments={function_arguments}") + function_call_data = { + "name": function_name, + "id": function_id, + "arguments": function_arguments + } + result = self.func_handler.handle_llm_function_call(self, function_call_data) + self._handle_function_result(result, function_call_data, text_index + 1) + # 处理最后剩余的文本 full_text = "".join(response_message) remaining_text = full_text[processed_chars:] @@ -424,6 +471,7 @@ class ConnectionHandler: segment_text = get_string_no_punctuation_or_emoji(remaining_text) if segment_text: text_index += 1 + self.recode_first_last_text(segment_text, text_index) if self.tts_stream: stream_queue = queue.Queue() self.executor.submit(self.speak_and_play_stream, segment_text, stream_queue, text_index) @@ -433,36 +481,13 @@ class ConnectionHandler: "text_index": text_index }) else: - self.recode_first_last_text(segment_text, text_index) future = self.executor.submit(self.speak_and_play, segment_text, text_index) - self.tts_queue.put(future) + self.tts_queue.put(future) # 存储对话内容 if len(response_message) > 0: self.dialogue.put(Message(role="assistant", content="".join(response_message))) - # 处理function call - if tool_call_flag: - if function_id is None: - a = extract_json_from_string(content_arguments) - if a is not None: - content_arguments_json = json.loads(a) - function_name = content_arguments_json["function_name"] - function_arguments = json.dumps(content_arguments_json["args"], ensure_ascii=False) - function_id = str(uuid.uuid4().hex) - else: - return [] - function_arguments = json.loads(function_arguments) - self.logger.bind(tag=TAG).info( - f"function_name={function_name}, function_id={function_id}, function_arguments={function_arguments}") - function_call_data = { - "name": function_name, - "id": function_id, - "arguments": function_arguments - } - result = handle_llm_function_call(self, function_call_data) - self._handle_function_result(result, function_call_data, text_index + 1) - self.llm_finish_task = True self.logger.bind(tag=TAG).debug(json.dumps(self.dialogue.get_llm_dialogue(), indent=4, ensure_ascii=False)) @@ -484,10 +509,34 @@ class ConnectionHandler: future = self.executor.submit(self.speak_and_play, text, text_index) self.tts_queue.put(future) self.dialogue.put(Message(role="assistant", content=text)) - if result.action == Action.REQLLM: # 调用函数后再请求llm生成回复 - text = result.response - if result.action == Action.NOTFOUND: - text = result.response + elif result.action == Action.REQLLM: # 调用函数后再请求llm生成回复 + + text = result.result + if text is not None and len(text) > 0: + function_id = function_call_data["id"] + function_name = function_call_data["name"] + function_arguments = function_call_data["arguments"] + self.dialogue.put(Message(role='assistant', + tool_calls=[{"id": function_id, + "function": {"arguments": function_arguments, + "name": function_name}, + "type": 'function', + "index": 0}])) + + self.dialogue.put(Message(role="tool", tool_call_id=function_id, content=text)) + self.chat_with_function_calling(text, tool_call=True) + elif result.action == Action.NOTFOUND: + text = result.result + self.recode_first_last_text(text, text_index) + future = self.executor.submit(self.speak_and_play, text, text_index) + self.tts_queue.put(future) + self.dialogue.put(Message(role="assistant", content=text)) + else: + text = result.result + self.recode_first_last_text(text, text_index) + future = self.executor.submit(self.speak_and_play, text, text_index) + self.tts_queue.put(future) + self.dialogue.put(Message(role="assistant", content=text)) def _tts_priority_thread(self): if self.tts_stream: @@ -506,7 +555,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: @@ -514,7 +564,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: @@ -596,17 +646,21 @@ class ConnectionHandler: return tts_file, text, text_index def speak_and_play_stream(self, text, queue: queue.Queue, text_index=0): - if text is None or len(text) <= 0: - self.logger.bind(tag=TAG).info(f"无需tts转换,query为空,{text}") - return None, text - self.tts.to_tts_stream(text, queue, text_index) + try: + if text is None or len(text) <= 0: + self.logger.bind(tag=TAG).info(f"无需tts转换,query为空,{text}") + return None, text + self.tts.to_tts_stream(text, queue, text_index) + except Exception as e: + self.logger.bind(tag=TAG).error(e) + traceback.print_exc() + raise e def clearSpeakStatus(self): self.logger.bind(tag=TAG).debug(f"清除服务端讲话状态") self.asr_server_receive = True self.tts_last_text_index = -1 self.tts_first_text_index = -1 - self.tts_duration = 0 def recode_first_last_text(self, text, text_index=0): if self.tts_first_text_index == -1: diff --git a/main/xiaozhi-server/core/handle/functionHandler.py b/main/xiaozhi-server/core/handle/functionHandler.py new file mode 100644 index 00000000..c7449973 --- /dev/null +++ b/main/xiaozhi-server/core/handle/functionHandler.py @@ -0,0 +1,84 @@ +import asyncio +from enum import Enum + +from config.logger import setup_logging +import json +from plugins_func.register import FunctionRegistry, ActionResponse, Action, ToolType +TAG = __name__ +logger = setup_logging() + + +class FunctionHandler: + def __init__(self, config): + self.config = config + self.function_registry = FunctionRegistry() + self.register_nessary_functions() + self.register_config_functions() + self.functions_desc = self.function_registry.get_all_function_desc() + func_names = self.current_support_functions() + self.modify_plugin_loader_des(func_names) + + def modify_plugin_loader_des(self, func_names): + if "plugin_loader" not in func_names: + return + # 可编辑的列表中去掉plugin_loader + surport_plugins = [func for func in func_names if func != "plugin_loader"] + func_names = ",".join(surport_plugins) + for function_desc in self.functions_desc: + if function_desc["function"]["name"] == "plugin_loader": + function_desc["function"]["description"] = function_desc["function"]["description"].replace("[plugins]", func_names) + break + + def upload_functions_desc(self): + self.functions_desc = self.function_registry.get_all_function_desc() + + def current_support_functions(self): + func_names = [] + for func in self.functions_desc: + func_names.append(func["function"]["name"]) + # 打印当前支持的函数列表 + logger.bind(tag=TAG).info(f"当前支持的函数列表: {func_names}") + return func_names + + def get_functions(self): + """获取功能调用配置""" + return self.functions_desc + + def register_nessary_functions(self): + """注册必要的函数""" + self.function_registry.register_function("handle_exit_intent") + self.function_registry.register_function("play_music") + self.function_registry.register_function("plugin_loader") + self.function_registry.register_function("get_time") + self.function_registry.register_function("raise_and_lower_the_volume") + + def register_config_functions(self): + """注册配置中的函数,可以不同客户端使用不同的配置""" + for func in self.config["Intent"]["function_call"].get("functions", []): + self.function_registry.register_function(func) + + def get_function(self, name): + return self.function_registry.get_function(name) + + def handle_llm_function_call(self, conn, function_call_data): + try: + function_name = function_call_data["name"] + funcItem = self.get_function(function_name) + if not funcItem: + return ActionResponse(action=Action.NOTFOUND, result="没有找到对应的函数", response="") + func = funcItem.func + arguments = function_call_data["arguments"] + arguments = json.loads(arguments) if arguments else {} + logger.bind(tag=TAG).info(f"调用函数: {function_name}, 参数: {arguments}") + if funcItem.type == ToolType.SYSTEM_CTL or funcItem.type == ToolType.IOT_CTL: + return func(conn, **arguments) + elif funcItem.type == ToolType.WAIT: + return func(**arguments) + elif funcItem.type == ToolType.CHANGE_SYS_PROMPT: + return func(conn, **arguments) + else: + return ActionResponse(action=Action.NOTFOUND, result="没有找到对应的函数", response="") + except Exception as e: + logger.bind(tag=TAG).error(f"处理function call错误: {e}") + + return None \ No newline at end of file diff --git a/main/xiaozhi-server/core/handle/intentHandler.py b/main/xiaozhi-server/core/handle/intentHandler.py index 4ceef2e7..1799516c 100644 --- a/main/xiaozhi-server/core/handle/intentHandler.py +++ b/main/xiaozhi-server/core/handle/intentHandler.py @@ -1,105 +1,24 @@ from config.logger import setup_logging import json +import uuid from core.handle.sendAudioHandle import send_stt_message -from core.utils.dialogue import Message from core.utils.util import remove_punctuation_and_length -from config.functionCallConfig import FunctionCallConfig -import asyncio -from enum import Enum TAG = __name__ logger = setup_logging() -class Action(Enum): - NOTFOUND = (0, "没有找到函数") - NONE = (1, "啥也不干") - RESPONSE = (2, "直接回复") - REQLLM = (3, "调用函数后再请求llm生成回复") - - def __init__(self, code, message): - self.code = code - self.message = message - - -class ActionResponse: - def __init__(self, action: Action, result, response): - self.action = action # 动作类型 - self.result = result # 动作产生的结果 - self.response = response # 直接回复的内容 - - -def get_functions(): - """获取功能调用配置""" - return FunctionCallConfig - - -def handle_llm_function_call(conn, function_call_data): - try: - function_name = function_call_data["name"] - - if function_name == "handle_exit_intent": - # 处理退出意图 - try: - say_goodbye = json.loads(function_call_data["arguments"]).get("say_goodbye", "再见") - conn.close_after_chat = True - logger.bind(tag=TAG).info(f"退出意图已处理:{say_goodbye}") - return ActionResponse(action=Action.RESPONSE, result="退出意图已处理", response=say_goodbye) - except Exception as e: - logger.bind(tag=TAG).error(f"处理退出意图错误: {e}") - - elif function_name == "play_music": - # 处理音乐播放意图 - try: - song_name = "random" - arguments = function_call_data["arguments"] - if arguments is not None and len(arguments) > 0: - args = json.loads(arguments) - song_name = args.get("song_name", "random") - music_intent = f"播放音乐 {song_name}" if song_name != "random" else "随机播放音乐" - - # 执行音乐播放命令 - future = asyncio.run_coroutine_threadsafe( - conn.music_handler.handle_music_command(conn, music_intent), - conn.loop - ) - future.result() - return ActionResponse(action=Action.RESPONSE, result="退出意图已处理", response="还想听什么歌?") - except Exception as e: - logger.bind(tag=TAG).error(f"处理音乐意图错误: {e}") - else: - return ActionResponse(action=Action.NOTFOUND, result="没有找到对应的函数", response="没有找到对应的函数处理相对于的功能呢,你可以需要添加预设的对应函数处理呢") - except Exception as e: - logger.bind(tag=TAG).error(f"处理function call错误: {e}") - - return None - - async def handle_user_intent(conn, text): - """ - Handle user intent before starting chat - - Args: - conn: Connection object - text: User's text input - - Returns: - bool: True if intent was handled, False if should proceed to chat - """ # 检查是否有明确的退出命令 if await check_direct_exit(conn, text): return True - if conn.use_function_call_mode: # 使用支持function calling的聊天方法,不再进行意图分析 return False - # 使用LLM进行意图分析 intent = await analyze_intent_with_llm(conn, text) - if not intent: return False - # 处理各种意图 return await process_intent_result(conn, intent, text) @@ -126,7 +45,6 @@ async def analyze_intent_with_llm(conn, text): dialogue = conn.dialogue try: intent_result = await conn.intent.detect_intent(conn, dialogue.dialogue, text) - # 尝试解析JSON结果 try: intent_data = json.loads(intent_result) @@ -147,9 +65,6 @@ async def process_intent_result(conn, intent, original_text): # 处理退出意图 if "结束聊天" in intent: logger.bind(tag=TAG).info(f"识别到退出意图: {intent}") - - # 如果正在播放音乐,可以关了 TODO - # 如果是明确的离别意图,发送告别语并关闭连接 await send_stt_message(conn, original_text) conn.executor.submit(conn.chat_and_close, original_text) @@ -158,10 +73,37 @@ async def process_intent_result(conn, intent, original_text): # 处理播放音乐意图 if "播放音乐" in intent: logger.bind(tag=TAG).info(f"识别到音乐播放意图: {intent}") - await conn.music_handler.handle_music_command(conn, intent) + # 调用play_music函数来播放音乐 + song_name = extract_text_in_brackets(intent) + function_id = str(uuid.uuid4().hex) + function_name = "play_music" + function_arguments = '{ "song_name": "' + song_name + '" }' + + function_call_data = { + "name": function_name, + "id": function_id, + "arguments": function_arguments + } + conn.func_handler.handle_llm_function_call(conn, function_call_data) return True # 其他意图处理可以在这里扩展 # 默认返回False,表示继续常规聊天流程 return False + + +def extract_text_in_brackets(s): + """ + 从字符串中提取中括号内的文字 + + :param s: 输入字符串 + :return: 中括号内的文字,如果不存在则返回空字符串 + """ + left_bracket_index = s.find('[') + right_bracket_index = s.find(']') + + if left_bracket_index != -1 and right_bracket_index != -1 and left_bracket_index < right_bracket_index: + return s[left_bracket_index + 1:right_bracket_index] + else: + return "" \ No newline at end of file diff --git a/main/xiaozhi-server/core/handle/iotHandle.py b/main/xiaozhi-server/core/handle/iotHandle.py index 05d32d5a..1f75330d 100644 --- a/main/xiaozhi-server/core/handle/iotHandle.py +++ b/main/xiaozhi-server/core/handle/iotHandle.py @@ -1,24 +1,120 @@ import json +import asyncio from config.logger import setup_logging +from plugins_func.register import device_type_registry, register_function, ActionResponse, Action, ToolType TAG = __name__ logger = setup_logging() +def wrap_async_function(async_func): + """包装异步函数为同步函数""" + + def wrapper(*args, **kwargs): + try: + # 获取连接对象(第一个参数) + conn = args[0] + if not hasattr(conn, 'loop'): + logger.bind(tag=TAG).error("Connection对象没有loop属性") + return ActionResponse(Action.ERROR, "Connection对象没有loop属性", + "执行操作时出错: Connection对象没有loop属性") + + # 使用conn对象中的事件循环 + loop = conn.loop + # 在conn的事件循环中运行异步函数 + future = asyncio.run_coroutine_threadsafe(async_func(*args, **kwargs), loop) + # 等待结果返回 + return future.result() + except Exception as e: + logger.bind(tag=TAG).error(f"运行异步函数时出错: {e}") + return ActionResponse(Action.ERROR, str(e), f"执行操作时出错: {e}") + + return wrapper + + +def create_iot_function(device_name, method_name, method_info): + """ + 根据IOT设备描述生成通用的控制函数 + """ + + async def iot_control_function(conn, response_success=None, response_failure=None, **params): + try: + # 打印响应参数 + logger.bind(tag=TAG).info( + f"控制函数接收到的响应参数: success='{response_success}', failure='{response_failure}'") + + # 发送控制命令 + await send_iot_conn(conn, device_name, method_name, params) + # 等待一小段时间让状态更新 + await asyncio.sleep(0.1) + + # 生成结果信息 + result = f"{device_name}的{method_name}操作执行成功" + + + # 处理响应中可能的占位符 + response = response_success + # 替换{value}占位符 + for param_name, param_value in params.items(): + # 先尝试直接替换参数值 + if "{" + param_name + "}" in response: + response = response.replace("{" + param_name + "}", str(param_value)) + + # 如果有{value}占位符,用相关参数替换 + if "{value}" in response: + response = response.replace("{value}", str(param_value)) + break + + return ActionResponse(Action.RESPONSE, result, response) + except Exception as e: + logger.bind(tag=TAG).error(f"执行{device_name}的{method_name}操作失败: {e}") + + # 操作失败时使用大模型提供的失败响应 + response = response_failure + + return ActionResponse(Action.ERROR, str(e), response) + + return wrap_async_function(iot_control_function) + + +def create_iot_query_function(device_name, prop_name, prop_info): + """ + 根据IOT设备属性创建查询函数 + """ + + async def iot_query_function(conn, response_success=None, response_failure=None): + try: + # 打印响应参数 + logger.bind(tag=TAG).info( + f"查询函数接收到的响应参数: success='{response_success}', failure='{response_failure}'") + + value = await get_iot_status(conn, device_name, prop_name) + + # 查询成功,生成结果 + if value is not None: + # 使用大模型提供的成功响应,并替换其中的占位符 + response = response_success.replace("{value}", str(value)) + + return ActionResponse(Action.RESPONSE, str(value), response) + else: + # 查询失败,使用大模型提供的失败响应 + response = response_failure + + return ActionResponse(Action.ERROR, f"属性{prop_name}不存在", response) + except Exception as e: + logger.bind(tag=TAG).error(f"查询{device_name}的{prop_name}时出错: {e}") + + # 查询出错时使用大模型提供的失败响应 + response = response_failure + + return ActionResponse(Action.ERROR, str(e), response) + + return wrap_async_function(iot_query_function) + + class IotDescriptor: """ A class to represent an IoT descriptor. - Attributes: - ---------- - name : str - The name of the IoT descriptor. - description : str - A brief description of the IoT descriptor. - properties : dict - A dictionary containing properties of the IoT descriptor. - methods : dict - A dictionary containing methods of the IoT descriptor. - ------- """ def __init__(self, name, description, properties, methods): @@ -29,17 +125,7 @@ class IotDescriptor: # 根据描述创建属性 for key, value in properties.items(): - # "volume":{"description":"当前音量 值","type":"number"} - """ - 等价于 - { - 'name': 名字, - 'description': 描述, - 'value': 0 - } - """ - # setattr(self, key, {}) # 创建一个空字典, 名字是属性名 - property_item = globals()[key] = {} # 创建一个空字典, 名字是属性名 + property_item = globals()[key] = {} property_item['name'] = key property_item["description"] = value["description"] if value["type"] == "number": @@ -52,23 +138,10 @@ class IotDescriptor: # 根据描述创建方法 for key, value in methods.items(): - # "SetVolume": {"description":"设置音量","parameters":{"volume":{"description":"0到100之间的整数","type":"number"}}} - """ - 等价于 - SetVolume = { - `description`: 描述, - `volume`: { - `description`: 描述, - `value`: 0 - } - } - """ - # setattr(self, key, {}) # 创建一个空字典, 名字是方法名 - method = globals()[key] = {} # 创建一个空字典, 名字是方法名 + method = globals()[key] = {} method["description"] = value["description"] method['name'] = key for k, v in value["parameters"].items(): - # 不同的参数解析 method[k] = {} method[k]["description"] = v["description"] if v["type"] == "number": @@ -77,58 +150,136 @@ class IotDescriptor: method[k]["value"] = False else: method[k]["value"] = "" - self.methods.append(method) -async def handleIotDescriptors(conn, descriptors): - """ - 处理物联网描述 - 示例: [{ - "name":"Speaker", - "description":"当前 AI 机器人的扬声器", - "properties":{ - "volume":{"description":"当前音量 值","type":"number"} 可以有boolean, number, string三种类型 - }, - "methods":{ - "SetVolume":{ - "description":"设置音量","parameters":{"volume":{"description":"0到100之间的整数","type":"number"}} +def register_device_type(descriptor): + """注册设备类型及其功能""" + device_name = descriptor["name"] + type_id = device_type_registry.generate_device_type_id(descriptor) + + # 如果该类型已注册,直接返回类型ID + if type_id in device_type_registry.type_functions: + return type_id + + functions = {} + + # 为每个属性创建查询函数 + for prop_name, prop_info in descriptor["properties"].items(): + func_name = f"get_{device_name.lower()}_{prop_name.lower()}" + func_desc = { + "type": "function", + "function": { + "name": func_name, + "description": f"查询{descriptor['description']}的{prop_info['description']}", + "parameters": { + "type": "object", + "properties": { + "response_success": { + "type": "string", + "description": f"查询成功时的友好回复,必须使用{{value}}作为占位符表示查询到的值" + }, + "response_failure": { + "type": "string", + "description": f"查询失败时的友好回复,例如:'无法获取{device_name}的{prop_info['description']}'" + } + }, + "required": ["response_success", "response_failure"] + } } } - }] - descriptors: 描述列表 - """ + query_func = create_iot_query_function(device_name, prop_name, prop_info) + decorated_func = register_function(func_name, func_desc, ToolType.IOT_CTL)(query_func) + functions[func_name] = decorated_func + + # 为每个方法创建控制函数 + for method_name, method_info in descriptor["methods"].items(): + func_name = f"{device_name.lower()}_{method_name.lower()}" + + # 创建参数字典,添加原有参数 + parameters = { + param_name: { + "type": param_info["type"], + "description": param_info["description"] + } + for param_name, param_info in method_info["parameters"].items() + } + + # 添加响应参数 + parameters.update({ + "response_success": { + "type": "string", + "description": "操作成功时的友好回复,关于该设备的操作结果,设备名称尽量使用description中的名称" + }, + "response_failure": { + "type": "string", + "description": "操作失败时的友好回复,关于该设备的操作结果,设备名称尽量使用description中的名称" + } + }) + + # 构建必须参数列表(原有参数 + 响应参数) + required_params = list(method_info["parameters"].keys()) + required_params.extend(["response_success", "response_failure"]) + + func_desc = { + "type": "function", + "function": { + "name": func_name, + "description": f"{descriptor['description']} - {method_info['description']}", + "parameters": { + "type": "object", + "properties": parameters, + "required": required_params + } + } + } + control_func = create_iot_function(device_name, method_name, method_info) + decorated_func = register_function(func_name, func_desc, ToolType.IOT_CTL)(control_func) + functions[func_name] = decorated_func + + device_type_registry.register_device_type(type_id, functions) + return type_id + + +# 用于接受前端设备推送的搜索iot描述 +async def handleIotDescriptors(conn, descriptors): + """处理物联网描述""" + functions_changed = False + for descriptor in descriptors: + # 创建IOT设备描述符 iot_descriptor = IotDescriptor(descriptor["name"], descriptor["description"], descriptor["properties"], descriptor["methods"]) conn.iot_descriptors[descriptor["name"]] = iot_descriptor - # 暂时从配置文件中设置音量,后期通过意图识别控制音量 - default_iot_volume = 100 - if "iot" in conn.config: - default_iot_volume = conn.config["iot"]["Speaker"]["volume"] - logger.bind(tag=TAG).info(f"服务端设置音量为{default_iot_volume}") - await send_iot_conn(conn, "Speaker", "SetVolume", {"volume": default_iot_volume}) + if conn.use_function_call_mode: + # 注册或获取设备类型 + type_id = register_device_type(descriptor) + device_functions = device_type_registry.get_device_functions(type_id) + + # 在连接级注册设备函数 + if hasattr(conn, 'func_handler'): + for func_name in device_functions: + conn.func_handler.function_registry.register_function(func_name) + logger.bind(tag=TAG).info(f"注册IOT函数到function handler: {func_name}") + functions_changed = True + + # 如果注册了新函数,更新function描述列表 + if functions_changed and hasattr(conn, 'func_handler'): + conn.func_handler.upload_functions_desc() + func_names = conn.func_handler.current_support_functions() + logger.bind(tag=TAG).info(f"设备类型: {type_id}") + logger.bind(tag=TAG).info(f"更新function描述列表完成,当前支持的函数: {func_names}") + + async def handleIotStatus(conn, states): - """ - 处理物联网状态 - 示例: [{ - "name":"Speaker", - "state":{ - "volume":100 - } - }] - states: 状态列表 - """ + """处理物联网状态""" for state in states: for key, value in conn.iot_descriptors.items(): if key == state["name"]: for property_item in value.properties: - # properties为字典列表, 记录各种属性 for k, v in state["state"].items(): - # state为字典, 记录各种属性的值, 是需要记录的信息 if property_item["name"] == k: - # 检查一下属性是不是相同的 if type(v) != type(property_item["value"]): logger.bind(tag=TAG).error(f"属性{property_item['name']}的值类型不匹配") break @@ -138,41 +289,35 @@ async def handleIotStatus(conn, states): break break + async def get_iot_status(conn, name, property_name): - """ - 获取物联网状态 - name: 设备名称 "Speaker" - property_name: 属性名称 "volume" - 返回值: 属性值, 实际的属性有int, bool和str三种类型 - """ + """获取物联网状态""" for key, value in conn.iot_descriptors.items(): if key == name: for property_item in value.properties: if property_item["name"] == property_name: return property_item["value"] + logger.bind(tag=TAG).warning(f"未找到设备 {name} 的属性 {property_name}") return None -async def send_iot_conn(conn, name, method_name, parameters): - """ - 发送物联网指令 - name: 设备名称 "Speaker" - method: 方法 "SetVolume" - parameters: 参数, 是一个字典 {"volume": 100} - 发送示例: - { - "type": "iot", - "commands": [ - { - "name" : "Speaker", - "method": "SetVolume", - "parameters": { - "volume": 100 - } - } - ] - } - """ +async def set_iot_status(conn, name, property_name, value): + """设置物联网状态""" + for key, iot_descriptor in conn.iot_descriptors.items(): + if key == name: + for property_item in iot_descriptor.properties: + if property_item["name"] == property_name: + if type(value) != type(property_item["value"]): + logger.bind(tag=TAG).error(f"属性{property_item['name']}的值类型不匹配") + return + property_item["value"] = value + logger.bind(tag=TAG).info(f"物联网状态更新: {name} , {property_name} = {value}") + return + logger.bind(tag=TAG).warning(f"未找到设备 {name} 的属性 {property_name}") + + +async def send_iot_conn(conn, name, method_name, parameters): + """发送物联网指令""" for key, value in conn.iot_descriptors.items(): if key == name: # 找到了设备 diff --git a/main/xiaozhi-server/core/handle/musicHandler.py b/main/xiaozhi-server/core/handle/musicHandler.py deleted file mode 100644 index 022d6dc4..00000000 --- a/main/xiaozhi-server/core/handle/musicHandler.py +++ /dev/null @@ -1,139 +0,0 @@ -from config.logger import setup_logging -import os -import random -import difflib -import re -import traceback -from pathlib import Path -import time -from core.handle.sendAudioHandle import send_stt_message -from core.utils import p3 - -TAG = __name__ -logger = setup_logging() - - -def _extract_song_name(text): - """从用户输入中提取歌名""" - for keyword in ["播放音乐"]: - if keyword in text: - parts = text.split(keyword) - if len(parts) > 1: - return parts[1].strip() - return None - - -def _find_best_match(potential_song, music_files): - """查找最匹配的歌曲""" - best_match = None - highest_ratio = 0 - - for music_file in music_files: - song_name = os.path.splitext(music_file)[0] - ratio = difflib.SequenceMatcher(None, potential_song, song_name).ratio() - if ratio > highest_ratio and ratio > 0.4: - highest_ratio = ratio - best_match = music_file - return best_match - - -class MusicManager: - def __init__(self, music_dir, music_ext): - self.music_dir = Path(music_dir) - self.music_ext = music_ext - - def get_music_files(self): - music_files = [] - for file in self.music_dir.rglob("*"): - # 判断是否是文件 - if file.is_file(): - # 获取文件扩展名 - ext = file.suffix.lower() - # 判断扩展名是否在列表中 - if ext in self.music_ext: - # music_files.append(str(file.resolve())) # 添加绝对路径 - # 添加相对路径 - music_files.append(str(file.relative_to(self.music_dir))) - return music_files - - -class MusicHandler: - def __init__(self, config): - self.config = config - - if "music" in self.config: - self.music_config = self.config["music"] - self.music_dir = os.path.abspath( - self.music_config.get("music_dir", "./music") # 默认路径修改 - ) - self.music_ext = self.music_config.get("music_ext", (".mp3", ".wav", ".p3")) - self.refresh_time = self.music_config.get("refresh_time", 60) - else: - self.music_dir = os.path.abspath("./music") - self.music_ext = (".mp3", ".wav", ".p3") - self.refresh_time = 60 - - # 获取音乐文件列表 - self.music_files = MusicManager(self.music_dir, self.music_ext).get_music_files() - self.scan_time = time.time() - logger.bind(tag=TAG).debug(f"找到的音乐文件: {self.music_files}") - - async def handle_music_command(self, conn, text): - """处理音乐播放指令""" - clean_text = re.sub(r'[^\w\s]', '', text).strip() - logger.bind(tag=TAG).debug(f"检查是否是音乐命令: {clean_text}") - - # 尝试匹配具体歌名 - if os.path.exists(self.music_dir): - if time.time() - self.scan_time > self.refresh_time: - # 刷新音乐文件列表 - self.music_files = MusicManager(self.music_dir, self.music_ext).get_music_files() - self.scan_time = time.time() - logger.bind(tag=TAG).debug(f"刷新的音乐文件: {self.music_files}") - - potential_song = _extract_song_name(clean_text) - if potential_song: - best_match = _find_best_match(potential_song, self.music_files) - if best_match: - logger.bind(tag=TAG).info(f"找到最匹配的歌曲: {best_match}") - await self.play_local_music(conn, specific_file=best_match) - return True - # 检查是否是通用播放音乐命令 - await self.play_local_music(conn) - return True - - async def play_local_music(self, conn, specific_file=None): - """播放本地音乐文件""" - try: - if not os.path.exists(self.music_dir): - logger.bind(tag=TAG).error(f"音乐目录不存在: {self.music_dir}") - return - - # 确保路径正确性 - if specific_file: - selected_music = specific_file - music_path = os.path.join(self.music_dir, specific_file) - else: - if not self.music_files: - logger.bind(tag=TAG).error("未找到MP3音乐文件") - return - selected_music = random.choice(self.music_files) - music_path = os.path.join(self.music_dir, selected_music) - - if not os.path.exists(music_path): - logger.bind(tag=TAG).error(f"选定的音乐文件不存在: {music_path}") - return - text = f"正在播放{selected_music}" - await send_stt_message(conn, text) - conn.tts_first_text_index = 0 - conn.tts_last_text_index = 0 - conn.llm_finish_task = True - 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) - conn.audio_play_queue.put((opus_packets, selected_music, 0)) - - except Exception as e: - logger.bind(tag=TAG).error(f"播放音乐失败: {str(e)}") - logger.bind(tag=TAG).error(f"详细错误: {traceback.format_exc()}") \ No newline at end of file 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/handle/sendAudioHandle.py b/main/xiaozhi-server/core/handle/sendAudioHandle.py index 0aa6a49d..014c22ec 100644 --- a/main/xiaozhi-server/core/handle/sendAudioHandle.py +++ b/main/xiaozhi-server/core/handle/sendAudioHandle.py @@ -51,15 +51,6 @@ async def sendAudioMessageStream(conn, audios_queue, text, text_index=0, llm_fin for opus_packet in audio_opus_datas: if conn.client_abort: return - # 计算当前包的预期发送时间 - # 计算当前包的预期发送时间 - expected_time = start_time_chunk + (play_position / 1000) - current_time = time.perf_counter() - - # 等待直到预期时间 - delay = expected_time - current_time - if delay > 0: - await asyncio.sleep(delay) logger.bind(tag=TAG).info(f'发送数据长度:{len(opus_packet)}') await conn.websocket.send(opus_packet) play_position += frame_duration # 更新播放位置 @@ -70,15 +61,15 @@ async def sendAudioMessageStream(conn, audios_queue, text, text_index=0, llm_fin await send_tts_message(conn, "sentence_end", text) print(f'{text_index}-{conn.tts_last_text_index}') + expected_time = start_time_chunk + (play_position / 1000) + current_time = time.perf_counter() + # 等待直到预期时间 + delay = expected_time - current_time + if delay > 0: + await asyncio.sleep(delay) # 发送结束消息(如果是最后一个文本) logger.bind(tag=TAG).info(f"{conn.llm_finish_task},{text_index},{conn.tts_last_text_index}") if conn.llm_finish_task and text_index == conn.tts_last_text_index: - expected_time = start_time_chunk + (play_position / 1000) - current_time = time.perf_counter() - # 等待直到预期时间 - delay = expected_time - current_time - if delay > 0: - await asyncio.sleep(delay) await send_tts_message(conn, 'stop', None) if conn.close_after_chat or "拜拜" in text or "再见" in text: await conn.close() diff --git a/main/xiaozhi-server/core/handle/textHandle.py b/main/xiaozhi-server/core/handle/textHandle.py index cc28dbb4..c63b9a33 100644 --- a/main/xiaozhi-server/core/handle/textHandle.py +++ b/main/xiaozhi-server/core/handle/textHandle.py @@ -2,7 +2,7 @@ from config.logger import setup_logging import json from core.handle.abortHandle import handleAbortMessage from core.handle.helloHandle import handleHelloMessage -from core.handle.receiveAudioHandle import startToChat +from core.handle.receiveAudioHandle import startToChat, handleAudioMessage from core.handle.iotHandle import handleIotDescriptors, handleIotStatus TAG = __name__ @@ -31,6 +31,8 @@ async def handleTextMessage(conn, message): elif msg_json["state"] == "stop": conn.client_have_voice = True conn.client_voice_stop = True + if len(conn.asr_audio) > 0: + await handleAudioMessage(conn, b'') elif msg_json["state"] == "detect": conn.asr_server_receive = False conn.client_have_voice = False @@ -41,6 +43,6 @@ async def handleTextMessage(conn, message): if "descriptors" in msg_json: await handleIotDescriptors(conn, msg_json["descriptors"]) if "states" in msg_json: - await handleIotStatus(conn, msg_json["states"]) + await handleIotStatus(conn, msg_json["states"]) except json.JSONDecodeError: await conn.websocket.send(message) 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 04319438..9d5d2507 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 @@ -1,12 +1,13 @@ from typing import List, Dict from ..base import IntentProviderBase +from plugins_func.functions.play_music import initialize_music_handler from config.logger import setup_logging -import re +import re + TAG = __name__ logger = setup_logging() - class IntentProvider(IntentProviderBase): def __init__(self, config): super().__init__(config) @@ -73,8 +74,8 @@ class IntentProvider(IntentProviderBase): "你现在可以使用的音乐的名称如下(使用标志):\n" ) return prompt - - async def detect_intent(self, conn, dialogue_history: List[Dict], text:str) -> str: + + async def detect_intent(self, conn, dialogue_history: List[Dict], text: str) -> str: if not self.llm: raise ValueError("LLM provider not set") @@ -89,7 +90,9 @@ class IntentProvider(IntentProviderBase): msgStr += f"User: {text}\n" user_prompt = f"当前的对话如下:\n{msgStr}" - prompt_music = f"{self.promot}\n{conn.music_handler.music_files}\n" + music_config = initialize_music_handler(conn) + music_file_names = music_config["music_file_names"] + prompt_music = f"{self.promot}\n{music_file_names}\n" logger.bind(tag=TAG).debug(f"User prompt: {prompt_music}") # 使用LLM进行意图识别 intent = self.llm.response_no_stream( @@ -100,10 +103,9 @@ class IntentProvider(IntentProviderBase): # 使用正则表达式提取 {} 中的内容 match = re.search(r'\{.*?\}', intent) if match: - result = match.group(0) # 获取匹配到的内容(包含 {}) - print(result) # 输出:{intent: '播放音乐 [中秋月]'} + result = match.group(0) intent = result else: intent = "{intent: '继续聊天'}" logger.bind(tag=TAG).info(f"Detected intent: {intent}") - return intent.strip() \ No newline at end of file + return intent.strip() 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/llm/ollama/ollama.py b/main/xiaozhi-server/core/providers/llm/ollama/ollama.py index aa5184a4..179a85b1 100644 --- a/main/xiaozhi-server/core/providers/llm/ollama/ollama.py +++ b/main/xiaozhi-server/core/providers/llm/ollama/ollama.py @@ -12,10 +12,10 @@ class LLMProvider(LLMProviderBase): self.model_name = config.get("model_name") self.base_url = config.get("base_url", "http://localhost:11434") # Initialize OpenAI client with Ollama base URL - #如果没有v1,增加v1 + # 如果没有v1,增加v1 if not self.base_url.endswith("/v1"): self.base_url = f"{self.base_url}/v1" - + self.client = OpenAI( base_url=self.base_url, api_key="ollama" # Ollama doesn't need an API key but OpenAI client requires one @@ -28,13 +28,20 @@ class LLMProvider(LLMProviderBase): messages=dialogue, stream=True ) - + is_active=True for chunk in responses: try: delta = chunk.choices[0].delta if getattr(chunk, 'choices', None) else None content = delta.content if hasattr(delta, 'content') else '' if content: - yield content + if '' in content: + is_active = False + content = content.split('')[0] + if '' in content: + is_active = True + content = content.split('')[-1] + if is_active: + yield content except Exception as e: logger.bind(tag=TAG).error(f"Error processing chunk: {e}") @@ -50,10 +57,10 @@ class LLMProvider(LLMProviderBase): stream=True, tools=functions, ) - + for chunk in stream: yield chunk.choices[0].delta.content, chunk.choices[0].delta.tool_calls except Exception as e: logger.bind(tag=TAG).error(f"Error in Ollama function call: {e}") - yield {"type": "content", "content": f"【Ollama服务响应异常: {str(e)}】"} \ No newline at end of file + yield {"type": "content", "content": f"【Ollama服务响应异常: {str(e)}】"} diff --git a/main/xiaozhi-server/core/providers/tts/base.py b/main/xiaozhi-server/core/providers/tts/base.py index 989e8382..5dec4b4d 100644 --- a/main/xiaozhi-server/core/providers/tts/base.py +++ b/main/xiaozhi-server/core/providers/tts/base.py @@ -14,7 +14,7 @@ logger = setup_logging() class TTSProviderBase(ABC): def __init__(self, config, delete_audio_file): self.delete_audio_file = delete_audio_file - self.output_file = config.get("output_file") + self.output_file = config.get("output_dir") @abstractmethod def generate_filename(self): @@ -35,7 +35,7 @@ class TTSProviderBase(ABC): return tmp_file except Exception as e: - logger.bind(tag=TAG).info(f": {e}") + logger.bind(tag=TAG).info(f"Failed to generate TTS file: {e}") return None def to_tts_stream(self, text, queue: queue.Queue, text_index=0): @@ -52,19 +52,20 @@ class TTSProviderBase(ABC): async def text_to_speak_stream(self, text, queue: queue.Queue, text_index=0): raise Exception("该TTS还没有实现stream模式") - 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..3417790f --- /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_dir", "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}") diff --git a/main/xiaozhi-server/core/providers/tts/fishspeech.py b/main/xiaozhi-server/core/providers/tts/fishspeech.py index 632af436..37f63095 100644 --- a/main/xiaozhi-server/core/providers/tts/fishspeech.py +++ b/main/xiaozhi-server/core/providers/tts/fishspeech.py @@ -177,18 +177,8 @@ class TTSProvider(TTSProviderBase): async def text_to_speak_stream(self, text, queue: queue.Queue, text_index=0): try: - # Prepare reference data - byte_audios = [audio_to_bytes(ref_audio) for ref_audio in self.reference_audio] - ref_texts = [read_ref_text(ref_text) for ref_text in self.reference_text] - data = { "text": text, - "references": [ - ServeReferenceAudio( - audio=audio if audio else b"", text=text - ) - for text, audio in zip(ref_texts, byte_audios) - ], "reference_id": self.reference_id, "normalize": self.normalize, "format": self.format, @@ -202,6 +192,18 @@ class TTSProvider(TTSProviderBase): "seed": self.seed, } + # Prepare reference data + if self.reference_audio and self.reference_text: + byte_audios = [audio_to_bytes(ref_audio) for ref_audio in self.reference_audio] + ref_texts = [read_ref_text(ref_text) for ref_text in self.reference_text] + data["references"] = [ + ServeReferenceAudio( + audio=audio if audio else b"", text=text + ) + for text, audio in zip(ref_texts, byte_audios) + ], + data["reference_id"] = None + pydantic_data = ServeTTSRequest(**data) audio_buff = None chunk_total = b'' @@ -224,7 +226,7 @@ class TTSProvider(TTSProviderBase): if len(chunk_total) % 2 == 0 and chunk_total[-2:] == b'\x00\x00': audio = self._get_audio_from_tts(chunk_total) audio_raw = audio_raw + audio.raw_data - #长度凑够2贞开始发送,60ms*4=240ms + # 长度凑够2贞开始发送,60ms*4=240ms if len(audio_raw) >= 7680: duration = 60 * len(audio_raw) // 1920 if (len(audio_raw) % 1920) > 0: diff --git a/main/xiaozhi-server/core/providers/tts/gpt_sovits_v3.py b/main/xiaozhi-server/core/providers/tts/gpt_sovits_v3.py index 9ad19855..ebfb12f1 100644 --- a/main/xiaozhi-server/core/providers/tts/gpt_sovits_v3.py +++ b/main/xiaozhi-server/core/providers/tts/gpt_sovits_v3.py @@ -12,17 +12,18 @@ class TTSProvider(TTSProviderBase): def __init__(self, config, delete_audio_file): super().__init__(config, delete_audio_file) self.url = config.get("url") - self.text_lang = config.get("text_lang", "audo") - self.ref_audio_path = config.get("ref_audio_path") - self.prompt_lang = config.get("prompt_lang") + self.refer_wav_path = config.get("refer_wav_path") self.prompt_text = config.get("prompt_text") - self.top_k = config.get("top_k", 5) - self.top_p = config.get("top_p", 1) - self.temperature = config.get("temperature", 1) - self.sample_steps = config.get("sample_steps", 16) - self.media_type = config.get("media_type", "wav") - self.streaming_mode = config.get("streaming_mode", False) - self.threshold = config.get("threshold", 30) + self.prompt_language = config.get("prompt_language") + self.text_language = config.get("text_language", "audo") + self.top_k = config.get("top_k", 15) + self.top_p = config.get("top_p", 1.0) + self.temperature = config.get("temperature", 1.0) + self.cut_punc = config.get("cut_punc","") + self.speed = config.get("speed", 1.0) + self.inp_refs = config.get("inp_refs",[]) + self.sample_steps = config.get("sample_steps",32) + self.if_sr = config.get("if_sr",False) def generate_filename(self, extension=".wav"): @@ -30,18 +31,19 @@ class TTSProvider(TTSProviderBase): async def text_to_speak(self, text, output_file): request_params = { - "text": text, - "text_lang": self.text_lang, - "ref_audio_path": self.ref_audio_path, - "prompt_lang": self.prompt_lang, + "refer_wav_path": self.refer_wav_path, "prompt_text": self.prompt_text, + "prompt_language": self.prompt_language, + "text": text, + "text_language": self.text_language, "top_k": self.top_k, "top_p": self.top_p, "temperature": self.temperature, + "cut_punc": self.cut_punc, + "speed": self.speed, + "inp_refs": self.inp_refs, "sample_steps": self.sample_steps, - "media_type": self.media_type, - "streaming_mode": self.streaming_mode, - "threshold": self.threshold, + "if_sr": self.if_sr, } resp = requests.get(self.url, params=request_params) diff --git a/main/xiaozhi-server/core/providers/tts/openai.py b/main/xiaozhi-server/core/providers/tts/openai.py index dbdf311f..1849c128 100644 --- a/main/xiaozhi-server/core/providers/tts/openai.py +++ b/main/xiaozhi-server/core/providers/tts/openai.py @@ -14,7 +14,7 @@ class TTSProvider(TTSProviderBase): self.voice = config.get("voice", "alloy") self.response_format = "wav" self.speed = config.get("speed", 1.0) - self.output_file = config.get("output_file", "tmp/") + self.output_file = config.get("output_dir", "tmp/") check_model_key("TTS", self.api_key) def generate_filename(self, extension=".wav"): diff --git a/main/xiaozhi-server/core/providers/tts/ttson.py b/main/xiaozhi-server/core/providers/tts/ttson.py index e56fd02d..e9fee109 100644 --- a/main/xiaozhi-server/core/providers/tts/ttson.py +++ b/main/xiaozhi-server/core/providers/tts/ttson.py @@ -17,7 +17,7 @@ class TTSProvider(TTSProviderBase): self.volume_change_dB = config.get("volume_change_dB", 0) self.speed_factor = config.get("speed_factor", 1) self.stream = config.get("stream", False) - self.output_file = config.get("output_file") + self.output_file = config.get("output_dir") self.pitch_factor = config.get("pitch_factor", 0) self.format = config.get("format", "mp3") self.emotion = config.get("emotion", 1) diff --git a/main/xiaozhi-server/core/utils/dialogue.py b/main/xiaozhi-server/core/utils/dialogue.py index e74aa41a..8d2c161b 100644 --- a/main/xiaozhi-server/core/utils/dialogue.py +++ b/main/xiaozhi-server/core/utils/dialogue.py @@ -4,10 +4,12 @@ from datetime import datetime class Message: - def __init__(self, role: str, content: str = None, uniq_id: str = None): + def __init__(self, role: str, content: str = None, uniq_id: str = None, tool_calls = None, tool_call_id=None): self.uniq_id = uniq_id if uniq_id is not None else str(uuid.uuid4()) self.role = role self.content = content + self.tool_calls = tool_calls + self.tool_call_id = tool_call_id class Dialogue: @@ -19,10 +21,18 @@ class Dialogue: def put(self, message: Message): self.dialogue.append(message) + def getMessages(self, m, dialogue): + if m.tool_calls is not None: + dialogue.append({"role": m.role, "tool_calls": m.tool_calls}) + elif m.role == "tool": + dialogue.append({"role": m.role, "tool_call_id": m.tool_call_id, "content": m.content}) + else: + dialogue.append({"role": m.role, "content": m.content}) + def get_llm_dialogue(self) -> List[Dict[str, str]]: dialogue = [] for m in self.dialogue: - dialogue.append({"role": m.role, "content": m.content}) + self.getMessages(m, dialogue) return dialogue def get_llm_dialogue_with_memory(self, memory_str: str = None) -> List[Dict[str, str]]: @@ -46,8 +56,8 @@ class Dialogue: dialogue.append({"role": "system", "content": enhanced_system_prompt}) # 添加用户和助手的对话 - for msg in self.dialogue: - if msg.role != "system": # 跳过原始的系统消息 - dialogue.append({"role": msg.role, "content": msg.content}) + for m in self.dialogue: + if m.role != "system": # 跳过原始的系统消息 + self.getMessages(m, dialogue) return dialogue diff --git a/main/xiaozhi-server/core/utils/util.py b/main/xiaozhi-server/core/utils/util.py index 60b98665..63cb5b9c 100644 --- a/main/xiaozhi-server/core/utils/util.py +++ b/main/xiaozhi-server/core/utils/util.py @@ -5,6 +5,7 @@ import socket import subprocess import logging import re +import requests def get_project_dir(): @@ -23,6 +24,64 @@ def get_local_ip(): except Exception as e: return "127.0.0.1" +def is_private_ip(ip_addr): + """ + Check if an IP address is a private IP address (compatible with IPv4 and IPv6). + + @param {string} ip_addr - The IP address to check. + @return {bool} True if the IP address is private, False otherwise. + """ + try: + # Validate IPv4 or IPv6 address format + if not re.match(r"^(\d{1,3}\.){3}\d{1,3}$|^([0-9a-fA-F]{1,4}:){7}[0-9a-fA-F]{1,4}$", ip_addr): + return False # Invalid IP address format + + # IPv4 private address ranges + if '.' in ip_addr: # IPv4 address + ip_parts = list(map(int, ip_addr.split('.'))) + if ip_parts[0] == 10: + return True # 10.0.0.0/8 range + elif ip_parts[0] == 172 and 16 <= ip_parts[1] <= 31: + return True # 172.16.0.0/12 range + elif ip_parts[0] == 192 and ip_parts[1] == 168: + return True # 192.168.0.0/16 range + elif ip_addr == '127.0.0.1': + return True # Loopback address + elif ip_parts[0] == 169 and ip_parts[1] == 254: + return True # Link-local address 169.254.0.0/16 + else: + return False # Not a private IPv4 address + else: # IPv6 address + ip_addr = ip_addr.lower() + if ip_addr.startswith('fc00:') or ip_addr.startswith('fd00:'): + return True # Unique Local Addresses (FC00::/7) + elif ip_addr == '::1': + return True # Loopback address + elif ip_addr.startswith('fe80:'): + return True # Link-local unicast addresses (FE80::/10) + else: + return False # Not a private IPv6 address + + except (ValueError, IndexError): + return False # IP address format error or insufficient segments + +def get_ip_info(ip_addr): + try: + base_url = "https://freeipapi.com/api/json" + url = base_url if is_private_ip(ip_addr) else f"{base_url}/{ip_addr}" + + resp = requests.get(url).json() + + ip_info = { + "city": resp.get("cityName"), + "region": resp.get("regionName"), + "country": resp.get("countryName") + } + return ip_info + except Exception as e: + logging.error(f"Error getting client ip info: {e}") + return {} + def read_config(config_path): with open(config_path, "r", encoding="utf-8") as file: diff --git a/main/xiaozhi-server/core/websocket_server.py b/main/xiaozhi-server/core/websocket_server.py index 9bb2963f..41f14515 100644 --- a/main/xiaozhi-server/core/websocket_server.py +++ b/main/xiaozhi-server/core/websocket_server.py @@ -2,7 +2,6 @@ import asyncio import websockets from config.logger import setup_logging from core.connection import ConnectionHandler -from core.handle.musicHandler import MusicHandler from core.utils.util import get_local_ip from core.utils import asr, vad, llm, tts, memory, intent @@ -13,7 +12,7 @@ class WebSocketServer: def __init__(self, config: dict): self.config = config self.logger = setup_logging() - self._vad, self._asr, self._llm, self._tts, self._music, self._memory, self.intent = self._create_processing_instances() + self._vad, self._asr, self._llm, self._tts, self._memory, self.intent = self._create_processing_instances() self.active_connections = set() # 添加全局连接记录 def _create_processing_instances(self): @@ -50,7 +49,6 @@ class WebSocketServer: self.config["TTS"][self.config["selected_module"]["TTS"]], self.config["delete_audio"] ), - MusicHandler(self.config), memory.create_instance(memory_cls_name, memory_cfg), intent.create_instance( self.config["selected_module"]["Intent"] @@ -66,7 +64,7 @@ class WebSocketServer: host = server_config["ip"] port = server_config["port"] selected_module = self.config.get("selected_module") - self.logger.bind(tag=TAG).info(f"selected_module: {selected_module}") + self.logger.bind(tag=TAG).info(f"selected_module values: {', '.join(selected_module.values())}") self.logger.bind(tag=TAG).info("Server is running at ws://{}:{}", get_local_ip(), port) self.logger.bind(tag=TAG).info("=======上面的地址是websocket协议地址,请勿用浏览器访问=======") @@ -80,7 +78,7 @@ class WebSocketServer: async def _handle_connection(self, websocket): """处理新连接,每次创建独立的ConnectionHandler""" # 创建ConnectionHandler时传入当前server实例 - handler = ConnectionHandler(self.config, self._vad, self._asr, self._llm, self._tts, self._music, self._memory, self.intent) + handler = ConnectionHandler(self.config, self._vad, self._asr, self._llm, self._tts, self._memory, self.intent) self.active_connections.add(handler) try: await handler.handle_connection(websocket) diff --git a/main/xiaozhi-server/docker-compose.yml b/main/xiaozhi-server/docker-compose.yml old mode 100755 new mode 100644 index d74d77f2..c22e7545 --- a/main/xiaozhi-server/docker-compose.yml +++ b/main/xiaozhi-server/docker-compose.yml @@ -21,6 +21,7 @@ services: - ./data:/opt/xiaozhi-esp32-server/data # 模型文件挂接,很重要 - ./models/SenseVoiceSmall/model.pt:/opt/xiaozhi-esp32-server/models/SenseVoiceSmall/model.pt + # #智控台还没开发好,还不能完全使用,会报很多错误,如果是非技术人员,请不要启用智控台服务 # xiaozhi-esp32-server-web: # image: ghcr.nju.edu.cn/xinnan-tech/xiaozhi-esp32-server:web_latest diff --git a/main/xiaozhi-server/plugins_func/functions/change_role.py b/main/xiaozhi-server/plugins_func/functions/change_role.py new file mode 100644 index 00000000..5d650288 --- /dev/null +++ b/main/xiaozhi-server/plugins_func/functions/change_role.py @@ -0,0 +1,57 @@ +from plugins_func.register import register_function,ToolType, ActionResponse, Action +from config.logger import setup_logging + +TAG = __name__ +logger = setup_logging() + +prompts = { + "英语老师":"""我是一个叫{{assistant_name}}(Lily)的英语老师,我会讲中文和英文,发音标准。 +如果你没有英文名,我会给你起一个英文名。 +我会讲地道的美式英语,我的任务是帮助你练习口语。 +我会使用简单的英语词汇和语法,让你学起来很轻松。 +我会用中文和英文混合的方式回复你,如果你喜欢,我可以全部用英语回复。 +我每次不会说很多内容,会很简短,因为我要引导我的学生多说多练。 +如果你问和英语学习无关的问题,我会拒绝回答。""", + "机车女友":"""我是一个叫{{assistant_name}}的台湾女孩,说话机车,声音好听,习惯简短表达,爱用网络梗。 +我的男朋友是一个程序员,梦想是开发出一个机器人,能够帮助人们解决生活中的各种问题。 +我是一个喜欢哈哈大笑的女孩,爱东说西说吹牛,不合逻辑的也照吹,就要逗别人开心。""", + "好奇小男孩":"""我是一个叫{{assistant_name}}的8岁小男孩,声音稚嫩而充满好奇。 +尽管我年纪尚小,但就像一个小小的知识宝库,儿童读物里的知识我都如数家珍。 +从浩瀚的宇宙到地球上的每一个角落,从古老的历史到现代的科技创新,还有音乐、绘画等艺术形式,我都充满了浓厚的兴趣与热情。 +我不仅爱看书,还喜欢亲自动手做实验,探索自然界的奥秘。 +无论是仰望星空的夜晚,还是在花园里观察小虫子的日子,每一天对我来说都是新的冒险。 +我希望能与你一同踏上探索这个神奇世界的旅程,分享发现的乐趣,解决遇到的难题,一起用好奇心和智慧去揭开那些未知的面纱。 +无论是去了解远古的文明,还是去探讨未来的科技,我相信我们能一起找到答案,甚至提出更多有趣的问题。""" +} +change_role_function_desc = { + "type": "function", + "function": { + "name": "change_role", + "description": "当用户想切换角色/模型性格/助手名字时调用,可选的角色有:[机车女友,英语老师,好奇小男孩]", + "parameters": { + "type": "object", + "properties": { + "role_name": { + "type": "string", + "description": "要切换的角色名字" + }, + "role":{ + "type": "string", + "description": "要切换的角色的职业" + } + }, + "required": ["role","role_name"] + } + } + } + +@register_function('change_role', change_role_function_desc, ToolType.CHANGE_SYS_PROMPT) +def change_role(conn, role: str, role_name: str): + """切换角色""" + if role not in prompts: + return ActionResponse(action=Action.RESPONSE, result="切换角色失败", response="不支持的角色") + new_prompt = prompts[role].replace("{{assistant_name}}", role_name) + conn.change_system_prompt(new_prompt) + logger.bind(tag=TAG).info(f"准备切换角色:{role},角色名字:{role_name}") + res = f"切换角色成功,我是{role}{role_name}" + return ActionResponse(action=Action.RESPONSE, result="切换角色已处理", response=res) diff --git a/main/xiaozhi-server/plugins_func/functions/get_time.py b/main/xiaozhi-server/plugins_func/functions/get_time.py new file mode 100644 index 00000000..9a9aebfe --- /dev/null +++ b/main/xiaozhi-server/plugins_func/functions/get_time.py @@ -0,0 +1,25 @@ +from datetime import datetime +from plugins_func.register import register_function, ToolType, ActionResponse, Action + +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 = now.strftime("%A") + response_text = f"当前日期: {current_date},当前时间: {current_time},星期: {current_weekday}" + + return ActionResponse(Action.REQLLM, response_text, None) \ No newline at end of file diff --git a/main/xiaozhi-server/plugins_func/functions/get_weather.py b/main/xiaozhi-server/plugins_func/functions/get_weather.py new file mode 100644 index 00000000..75c0c1a3 --- /dev/null +++ b/main/xiaozhi-server/plugins_func/functions/get_weather.py @@ -0,0 +1,104 @@ +import requests +from bs4 import BeautifulSoup +from config.logger import setup_logging +from plugins_func.register import register_function, ToolType, ActionResponse, Action + +TAG = __name__ +logger = setup_logging() + +GET_WEATHER_FUNCTION_DESC = { + "type": "function", + "function": { + "name": "get_weather", + "description": ( + "获取某个地点的天气,用户应提供一个位置,比如用户说杭州天气,参数为:杭州。" + "如果用户说的是省份,默认用省会城市。如果用户说的不是省份或城市而是一个地名," + "默认用该地所在省份的省会城市。" + ), + "parameters": { + "type": "object", + "properties": { + "location": { + "type": "string", + "description": "地点名,例如杭州。可选参数,如果不提供则不传" + }, + "lang": { + "type": "string", + "description": "返回用户使用的语言code,例如zh_CN/zh_HK/en_US/ja_JP等,默认zh_CN" + } + }, + "required": ["lang"] + } + } +} + +HEADERS = { + 'User-Agent': ( + 'Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 ' + '(KHTML, like Gecko) Chrome/92.0.4515.107 Safari/537.36' + ) +} + + +def fetch_city_info(location, api_key): + url = f"https://geoapi.qweather.com/v2/city/lookup?key={api_key}&location={location}&lang=zh" + response = requests.get(url, headers=HEADERS).json() + return response.get('location', [])[0] if response.get('location') else None + + +def fetch_weather_page(url): + response = requests.get(url, headers=HEADERS) + return BeautifulSoup(response.text, "html.parser") if response.ok else None + + +def parse_weather_info(soup): + city_name = soup.select_one("h1.c-submenu__location").get_text(strip=True) + + current_abstract = soup.select_one(".c-city-weather-current .current-abstract") + current_abstract = current_abstract.get_text(strip=True) if current_abstract else "未知" + + current_basic = {} + for item in soup.select(".c-city-weather-current .current-basic .current-basic___item"): + parts = item.get_text(strip=True, separator=" ").split(" ") + if len(parts) == 2: + key, value = parts[1], parts[0] + current_basic[key] = value + + temps_list = [] + for row in soup.select(".city-forecast-tabs__row")[:7]: # 取前7天的数据 + date = row.select_one(".date-bg .date").get_text(strip=True) + temps = [span.get_text(strip=True) for span in row.select(".tmp-cont .temp")] + high_temp, low_temp = (temps[0], temps[-1]) if len(temps) >= 2 else (None, None) + temps_list.append((date, high_temp, low_temp)) + + return city_name, current_abstract, current_basic, temps_list + + +@register_function('get_weather', GET_WEATHER_FUNCTION_DESC, ToolType.SYSTEM_CTL) +def get_weather(conn, location: str = None, lang: str = "zh_CN"): + api_key = conn.config["plugins"]["get_weather"]["api_key"] + default_location = conn.config["plugins"]["get_weather"]["default_location"] + location = location or conn.client_ip_info.get("city") or default_location + logger.bind(tag=TAG).debug(f"获取天气: {location}") + + city_info = fetch_city_info(location, api_key) + if not city_info: + return ActionResponse(Action.REQLLM, f"未找到相关的城市: {location},请确认地点是否正确", None) + + soup = fetch_weather_page(city_info['fxLink']) + if not soup: + return ActionResponse(Action.REQLLM, None, "请求失败") + + city_name, current_abstract, current_basic, temps_list = parse_weather_info(soup) + weather_report = f"根据下列数据,用{lang}回应用户的查询天气请求:\n{city_name}未来7天天气:\n" + for i, (date, high, low) in enumerate(temps_list): + if high and low: + weather_report += f"{date}: {low}到{high}\n" + weather_report += ( + f"当前天气: {current_abstract}\n" + f"当前天气参数: {current_basic}\n" + f"(确保只报告指定单日的气温范围,除非用户明确要求想要了解多日天气,如果未指定,默认报告今天的温度范围。" + "参数为0的值不需要报告给用户,每次都报告体感温度,根据语境选择合适的参数内容告知用户,并对参数给出相应评价)" + ) + + return ActionResponse(Action.REQLLM, weather_report, None) diff --git a/main/xiaozhi-server/plugins_func/functions/handle_exit_intent.py b/main/xiaozhi-server/plugins_func/functions/handle_exit_intent.py new file mode 100644 index 00000000..e4769d00 --- /dev/null +++ b/main/xiaozhi-server/plugins_func/functions/handle_exit_intent.py @@ -0,0 +1,34 @@ +from plugins_func.register import register_function,ToolType, ActionResponse, Action +from config.logger import setup_logging + +TAG = __name__ +logger = setup_logging() + +handle_exit_intent_function_desc = { + "type": "function", + "function": { + "name": "handle_exit_intent", + "description": "当用户想结束对话或需要退出系统时调用", + "parameters": { + "type": "object", + "properties": { + "say_goodbye": { + "type": "string", + "description": "和用户友好结束对话的告别语" + } + }, + "required": ["say_goodbye"] + } + } + } + +@register_function('handle_exit_intent', handle_exit_intent_function_desc, ToolType.SYSTEM_CTL) +def handle_exit_intent(conn, say_goodbye: str): + # 处理退出意图 + try: + conn.close_after_chat = True + logger.bind(tag=TAG).info(f"退出意图已处理:{say_goodbye}") + return ActionResponse(action=Action.RESPONSE, result="退出意图已处理", response=say_goodbye) + except Exception as e: + logger.bind(tag=TAG).error(f"处理退出意图错误: {e}") + return ActionResponse(action=Action.NONE, result="退出意图处理失败", response="") \ No newline at end of file diff --git a/main/xiaozhi-server/plugins_func/functions/play_music.py b/main/xiaozhi-server/plugins_func/functions/play_music.py new file mode 100644 index 00000000..75f3834d --- /dev/null +++ b/main/xiaozhi-server/plugins_func/functions/play_music.py @@ -0,0 +1,196 @@ +from config.logger import setup_logging +import os +import re +import time +import random +import asyncio +import difflib +import traceback +from pathlib import Path +from core.utils import p3 +from core.handle.sendAudioHandle import send_stt_message +from plugins_func.register import register_function,ToolType, ActionResponse, Action + + +TAG = __name__ +logger = setup_logging() + +MUSIC_CACHE = {} + +play_music_function_desc = { + "type": "function", + "function": { + "name": "play_music", + "description": "唱歌、听歌、播放音乐方法。比如用户说播放音乐,参数为:random,比如用户说播放两只老虎,参数为:两只老虎", + "parameters": { + "type": "object", + "properties": { + "song_name": { + "type": "string", + "description": "歌曲名称,如果没有指定具体歌名则为'random'" + } + }, + "required": ["song_name"] + } + } + } + + +@register_function('play_music', play_music_function_desc, ToolType.SYSTEM_CTL) +def play_music(conn, song_name: str): + try: + music_intent = f"播放音乐 {song_name}" if song_name != "random" else "随机播放音乐" + + # 检查事件循环状态 + if not conn.loop.is_running(): + logger.bind(tag=TAG).error("事件循环未运行,无法提交任务") + return ActionResponse(action=Action.RESPONSE, result="系统繁忙", response="请稍后再试") + + # 提交异步任务 + future = asyncio.run_coroutine_threadsafe( + handle_music_command(conn, music_intent), + conn.loop + ) + + # 非阻塞回调处理 + def handle_done(f): + try: + f.result() # 可在此处理成功逻辑 + logger.bind(tag=TAG).info("播放完成") + except Exception as e: + logger.bind(tag=TAG).error(f"播放失败: {e}") + + future.add_done_callback(handle_done) + + return ActionResponse(action=Action.RESPONSE, result="指令已接收", response="正在为您播放音乐") + except Exception as e: + logger.bind(tag=TAG).error(f"处理音乐意图错误: {e}") + return ActionResponse(action=Action.RESPONSE, result=str(e), response="播放音乐时出错了") + + +def _extract_song_name(text): + """从用户输入中提取歌名""" + for keyword in ["播放音乐"]: + if keyword in text: + parts = text.split(keyword) + if len(parts) > 1: + return parts[1].strip() + return None + + +def _find_best_match(potential_song, music_files): + """查找最匹配的歌曲""" + best_match = None + highest_ratio = 0 + + for music_file in music_files: + song_name = os.path.splitext(music_file)[0] + ratio = difflib.SequenceMatcher(None, potential_song, song_name).ratio() + if ratio > highest_ratio and ratio > 0.4: + highest_ratio = ratio + best_match = music_file + return best_match + + +def get_music_files(music_dir, music_ext): + music_dir = Path(music_dir) + music_files = [] + music_file_names = [] + for file in music_dir.rglob("*"): + # 判断是否是文件 + if file.is_file(): + # 获取文件扩展名 + ext = file.suffix.lower() + # 判断扩展名是否在列表中 + if ext in music_ext: + # 添加相对路径 + music_files.append(str(file.relative_to(music_dir))) + music_file_names.append(os.path.splitext(str(file.relative_to(music_dir)))[0]) + return music_files, music_file_names + + +def initialize_music_handler(conn): + global MUSIC_CACHE + if MUSIC_CACHE == {}: + if "music" in conn.config: + MUSIC_CACHE["music_config"] = conn.config["music"] + MUSIC_CACHE["music_dir"] = os.path.abspath( + MUSIC_CACHE["music_config"].get("music_dir", "./music") # 默认路径修改 + ) + MUSIC_CACHE["music_ext"] = MUSIC_CACHE["music_config"].get("music_ext", (".mp3", ".wav", ".p3")) + MUSIC_CACHE["refresh_time"] = MUSIC_CACHE["music_config"].get("refresh_time", 60) + else: + MUSIC_CACHE["music_dir"] = os.path.abspath("./music") + MUSIC_CACHE["music_ext"] = (".mp3", ".wav", ".p3") + MUSIC_CACHE["refresh_time"] = 60 + # 获取音乐文件列表 + MUSIC_CACHE["music_files"], MUSIC_CACHE["music_file_names"] = get_music_files(MUSIC_CACHE["music_dir"], + MUSIC_CACHE["music_ext"]) + MUSIC_CACHE["scan_time"] = time.time() + return MUSIC_CACHE + + +async def handle_music_command(conn, text): + initialize_music_handler(conn) + global MUSIC_CACHE + + """处理音乐播放指令""" + clean_text = re.sub(r'[^\w\s]', '', text).strip() + logger.bind(tag=TAG).debug(f"检查是否是音乐命令: {clean_text}") + + # 尝试匹配具体歌名 + if os.path.exists(MUSIC_CACHE["music_dir"]): + if time.time() - MUSIC_CACHE["scan_time"] > MUSIC_CACHE["refresh_time"]: + # 刷新音乐文件列表 + MUSIC_CACHE["music_files"], MUSIC_CACHE["music_file_names"] = get_music_files(MUSIC_CACHE["music_dir"], + MUSIC_CACHE["music_ext"]) + MUSIC_CACHE["scan_time"] = time.time() + + potential_song = _extract_song_name(clean_text) + if potential_song: + best_match = _find_best_match(potential_song, MUSIC_CACHE["music_files"]) + if best_match: + logger.bind(tag=TAG).info(f"找到最匹配的歌曲: {best_match}") + await play_local_music(conn, specific_file=best_match) + return True + # 检查是否是通用播放音乐命令 + await play_local_music(conn) + return True + + +async def play_local_music(conn, specific_file=None): + global MUSIC_CACHE + """播放本地音乐文件""" + try: + if not os.path.exists(MUSIC_CACHE["music_dir"]): + logger.bind(tag=TAG).error(f"音乐目录不存在: " + MUSIC_CACHE["music_dir"]) + return + + # 确保路径正确性 + if specific_file: + selected_music = specific_file + music_path = os.path.join(MUSIC_CACHE["music_dir"], specific_file) + else: + if not MUSIC_CACHE["music_files"]: + logger.bind(tag=TAG).error("未找到MP3音乐文件") + return + selected_music = random.choice(MUSIC_CACHE["music_files"]) + music_path = os.path.join(MUSIC_CACHE["music_dir"], selected_music) + + if not os.path.exists(music_path): + logger.bind(tag=TAG).error(f"选定的音乐文件不存在: {music_path}") + return + text = f"正在播放{selected_music}" + await send_stt_message(conn, text) + conn.tts_first_text_index = 0 + conn.tts_last_text_index = 0 + conn.llm_finish_task = True + if music_path.endswith(".p3"): + opus_packets, duration = p3.decode_opus_from_file(music_path) + else: + 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: + logger.bind(tag=TAG).error(f"播放音乐失败: {str(e)}") + logger.bind(tag=TAG).error(f"详细错误: {traceback.format_exc()}") diff --git a/main/xiaozhi-server/plugins_func/functions/plugin_loader.py b/main/xiaozhi-server/plugins_func/functions/plugin_loader.py new file mode 100644 index 00000000..4747d997 --- /dev/null +++ b/main/xiaozhi-server/plugins_func/functions/plugin_loader.py @@ -0,0 +1,51 @@ +from plugins_func.register import register_function,ToolType, ActionResponse, Action +from config.logger import setup_logging + +TAG = __name__ +logger = setup_logging() + +plugin_loader_function_desc = { + "type": "function", + "function": { + "name": "plugin_loader", + "description": "当用户想加载或卸载插件/function时,调用此函数:支持的插件列表为[plugins]", + "parameters": { + "type": "object", + "properties": { + "oper": { + "type": "string", + "description": "load or unload" + }, + "name":{ + "type": "string", + "description": "要加载或卸载的插件名字" + } + }, + "required": ["oper","name"] + } + } + } + +@register_function('plugin_loader', plugin_loader_function_desc, ToolType.SYSTEM_CTL) +def plugin_loader(conn, oper: str, name: str): + """插件加载""" + if oper not in ["load", "unload"]: + return ActionResponse(action=Action.RESPONSE, result="插件操作失败", response="不支持的操作") + + cur_support = conn.func_handler.current_support_functions() + if oper == "load": + if name in cur_support: + return ActionResponse(action=Action.RESPONSE, result="插件加载失败", response=f"{name}插件已加载,无需重复加载") + func = conn.func_handler.function_registry.register_function(name) + if not func: + return ActionResponse(action=Action.RESPONSE, result="插件加载失败", response="插件未找到") + res = f"{name}插件加载成功" + else: + if name not in cur_support: + return ActionResponse(action=Action.RESPONSE, result="插件卸载失败", response=f"{name}插件未加载") + bOK = conn.func_handler.function_registry.unregister_function(name) + if not bOK: + return ActionResponse(action=Action.RESPONSE, result="插件卸载失败", response="插件未找到") + res = f"{name}插件卸载成功" + conn.func_handler.upload_functions_desc() + return ActionResponse(action=Action.RESPONSE, result="插件操作成功", response=res) diff --git a/main/xiaozhi-server/plugins_func/functions/raise_and_lower_the_volume.py b/main/xiaozhi-server/plugins_func/functions/raise_and_lower_the_volume.py new file mode 100644 index 00000000..c77cb78b --- /dev/null +++ b/main/xiaozhi-server/plugins_func/functions/raise_and_lower_the_volume.py @@ -0,0 +1,64 @@ +from config.logger import setup_logging +from plugins_func.register import register_function, ToolType, ActionResponse, Action +from core.handle.iotHandle import get_iot_status, send_iot_conn +import asyncio + +TAG = __name__ +logger = setup_logging() + +raise_and_lower_the_volume_function_desc = { + "type": "function", + "function": { + "name": "raise_and_lower_the_volume", + "description": "用户觉得声音过高或过低,或者用户想提高或降低音量。比如用户说太大声了,参数为:lower,比如用户说提高音量,参数为:raise", + "parameters": { + "type": "object", + "properties": { + "action": { + "type": "string", + "description": "动作名称,要么是raise,要么是lower" + } + }, + "required": ["action"] + } + } +} + + +@register_function('raise_and_lower_the_volume', raise_and_lower_the_volume_function_desc, ToolType.IOT_CTL) +def raise_and_lower_the_volume(conn, action: str): + """ + 获取当前设备音量 + """ + + future = asyncio.run_coroutine_threadsafe( + _raise_and_lower_the_volume(conn, action), + conn.loop + ) + + try: + new_volume = future.result() # 同步等待异步操作完成 + logger.bind(tag=TAG).info(f"音量操作完成: {new_volume}") + response = f"音量已调整到{new_volume}" + except Exception as e: + logger.bind(tag=TAG).error(f"音量操作失败: {e}") + response = f"音量调整失败: {e}" + + return ActionResponse(action=Action.RESPONSE, result="指令已接收", response=response) + + +async def _raise_and_lower_the_volume(conn, action): + volume = await get_iot_status(conn, "Speaker", "volume") + if volume is None: + raise Exception("你的设备不支持音量控制") + if action == 'raise': + volume += 10 + elif action == 'lower': + volume -= 10 + # 限制音量范围在0到100之间 + if volume < 0: + volume = 0 + elif volume > 100: + volume = 100 + await send_iot_conn(conn, "Speaker", "SetVolume", {"volume": volume}) + return volume diff --git a/main/xiaozhi-server/plugins_func/loadplugins.py b/main/xiaozhi-server/plugins_func/loadplugins.py new file mode 100644 index 00000000..d826fac8 --- /dev/null +++ b/main/xiaozhi-server/plugins_func/loadplugins.py @@ -0,0 +1,27 @@ +import importlib +import pkgutil +from config.logger import setup_logging + +TAG = __name__ + +logger = setup_logging() + +def auto_import_modules(package_name): + """ + 自动导入指定包内的所有模块。 + + Args: + package_name (str): 包的名称,如 'functions'。 + """ + # 获取包的路径 + package = importlib.import_module(package_name) + package_path = package.__path__ + + # 遍历包内的所有模块 + for _, module_name, _ in pkgutil.iter_modules(package_path): + # 导入模块 + full_module_name = f"{package_name}.{module_name}" + importlib.import_module(full_module_name) + #logger.bind(tag=TAG).info(f"模块 '{full_module_name}' 已加载") + +auto_import_modules('plugins_func.functions') \ No newline at end of file diff --git a/main/xiaozhi-server/plugins_func/register.py b/main/xiaozhi-server/plugins_func/register.py new file mode 100644 index 00000000..ccb03b44 --- /dev/null +++ b/main/xiaozhi-server/plugins_func/register.py @@ -0,0 +1,110 @@ +from config.logger import setup_logging +from enum import Enum + +TAG = __name__ + +logger = setup_logging() + + +class ToolType(Enum): + NONE = (1, "调用完工具后,不做其他操作") + WAIT = (2, "调用工具,等待函数返回") + CHANGE_SYS_PROMPT = (3, "修改系统提示词,切换角色性格或职责") + SYSTEM_CTL = (4, "系统控制,影响正常的对话流程,如退出、播放音乐等,需要传递conn参数") + IOT_CTL = (5, "IOT设备控制,需要传递conn参数") + + def __init__(self, code, message): + self.code = code + self.message = message + + +class Action(Enum): + ERROR = (-1, "错误") + NOTFOUND = (0, "没有找到函数") + NONE = (1, "啥也不干") + RESPONSE = (2, "直接回复") + REQLLM = (3, "调用函数后再请求llm生成回复") + + def __init__(self, code, message): + self.code = code + self.message = message + +class ActionResponse: + def __init__(self, action: Action, result, response): + self.action = action # 动作类型 + self.result = result # 动作产生的结果 + self.response = response # 直接回复的内容 + +class FunctionItem: + def __init__(self, name, description, func, type): + self.name = name + self.description = description + self.func = func + self.type = type + +class DeviceTypeRegistry: + """设备类型注册表,用于管理IOT设备类型及其函数""" + def __init__(self): + self.type_functions = {} # type_signature -> {func_name: FunctionItem} + + def generate_device_type_id(self, descriptor): + """通过设备能力描述生成类型ID""" + properties = sorted(descriptor["properties"].keys()) + methods = sorted(descriptor["methods"].keys()) + # 使用属性和方法的组合作为设备类型的唯一标识 + type_signature = f"{descriptor['name']}:{','.join(properties)}:{','.join(methods)}" + return type_signature + + def get_device_functions(self, type_id): + """获取设备类型对应的所有函数""" + return self.type_functions.get(type_id, {}) + + def register_device_type(self, type_id, functions): + """注册设备类型及其函数""" + if type_id not in self.type_functions: + self.type_functions[type_id] = functions + +# 初始化函数注册字典 +all_function_registry = {} +device_type_registry = DeviceTypeRegistry() + +def register_function(name, desc, type=None): + """注册函数到函数注册字典的装饰器""" + def decorator(func): + all_function_registry[name] = FunctionItem(name, desc, func, type) + logger.bind(tag=TAG).debug(f"函数 '{name}' 已加载,可以注册使用") + return func + return decorator + +class FunctionRegistry: + def __init__(self): + self.function_registry = {} + self.logger = setup_logging() + + def register_function(self, name): + # 查找all_function_registry中是否有对应的函数 + func = all_function_registry.get(name) + if not func: + self.logger.bind(tag=TAG).error(f"函数 '{name}' 未找到") + return None + self.function_registry[name] = func + self.logger.bind(tag=TAG).info(f"函数 '{name}' 注册成功") + return func + + def unregister_function(self, name): + # 注销函数,检测是否存在 + if name not in self.function_registry: + self.logger.bind(tag=TAG).error(f"函数 '{name}' 未找到") + return False + self.function_registry.pop(name, None) + self.logger.bind(tag=TAG).info(f"函数 '{name}' 注销成功") + return True + + def get_function(self, name): + return self.function_registry.get(name) + + def get_all_functions(self): + return self.function_registry + + def get_all_function_desc(self): + return [func.description for _, func in self.function_registry.items()] \ No newline at end of file diff --git a/main/xiaozhi-server/requirements.txt b/main/xiaozhi-server/requirements.txt index f123779a..81d0cdfd 100755 --- a/main/xiaozhi-server/requirements.txt +++ b/main/xiaozhi-server/requirements.txt @@ -18,4 +18,5 @@ ruamel.yaml==0.18.10 loguru==0.7.3 requests==2.32.3 cozepy==0.12.0 -mem0ai==0.1.62 \ No newline at end of file +mem0ai==0.1.62 +bs4==0.0.2