diff --git a/README.md b/README.md index 943501a7..eb718c24 100644 --- a/README.md +++ b/README.md @@ -183,8 +183,8 @@ Spearheaded by Professor Siyuan Liu's Team (South China University of Technology #### 🚀 部署方式选择 | 部署方式 | 特点 | 适用场景 | 部署文档 | 配置要求 | 视频教程 | |---------|------|---------|---------|---------|---------| -| **最简化安装** | 智能对话、IOT、MCP、视觉感知 | 低配置环境,数据存储在配置文件,无需数据库 | [①Docker版](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E5%8F%AA%E8%BF%90%E8%A1%8Cserver) / [②源码部署](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E5%8F%AA%E8%BF%90%E8%A1%8Cserver)| 如果使用`FunASR`要2核4G,如果全API,要2核2G | - | -| **全模块安装** | 智能对话、IOT、MCP接入点、声纹识别、视觉感知、OTA、智控台 | 完整功能体验,数据存储在数据库 |[①Docker版](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) / [②源码部署](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) / [③源码部署自动更新教程](./docs/dev-ops-integration.md) | 如果使用`FunASR`要4核8G,如果全API,要2核4G| [本地源码启动视频教程](https://www.bilibili.com/video/BV1wBJhz4Ewe) | +| **最简化安装** | 智能对话、单智能体管理 | 低配置环境,数据存储在配置文件,无需数据库 | [①Docker版](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E5%8F%AA%E8%BF%90%E8%A1%8Cserver) / [②源码部署](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E5%8F%AA%E8%BF%90%E8%A1%8Cserver)| 如果使用`FunASR`要2核4G,如果全API,要2核2G | - | +| **全模块安装** | 智能对话、多用户管理、多智能体管理、智控台界面操作 | 完整功能体验,数据存储在数据库 |[①Docker版](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) / [②源码部署](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) / [③源码部署自动更新教程](./docs/dev-ops-integration.md) | 如果使用`FunASR`要4核8G,如果全API,要2核4G| [本地源码启动视频教程](https://www.bilibili.com/video/BV1wBJhz4Ewe) | 常见问题及相关教程,可参考[这个链接](./docs/FAQ.md) @@ -211,10 +211,10 @@ Websocket接口地址: wss://2662r3426b.vicp.fun/xiaozhi/v1/ | 模块名称 | 入门全免费设置 | 流式配置 | |:---:|:---:|:---:| -| ASR(语音识别) | FunASR(本地) | 👍FunASR(本地GPU模式) | -| LLM(大模型) | ChatGLMLLM(智谱glm-4-flash) | 👍AliLLM(qwen3-235b-a22b-instruct-2507) 或 👍DoubaoLLM(doubao-1-5-pro-32k-250115) | -| VLLM(视觉大模型) | ChatGLMVLLM(智谱glm-4v-flash) | 👍QwenVLVLLM(千问qwen2.5-vl-3b-instructh) | -| TTS(语音合成) | ✅LinkeraiTTS(灵犀流式) | 👍HuoshanDoubleStreamTTS(火山双流式语音合成) 或 👍AliyunStreamTTS(阿里云流式语音合成) | +| ASR(语音识别) | FunASR(本地) | 👍XunfeiStreamASR(讯飞流式) | +| LLM(大模型) | glm-4-flash(智谱) | 👍qwen-flash(阿里百炼) | +| VLLM(视觉大模型) | glm-4v-flash(智谱) | 👍qwen2.5-vl-3b-instructh(阿里百炼) | +| TTS(语音合成) | ✅LinkeraiTTS(灵犀流式) | 👍HuoshanDoubleStreamTTS(火山流式) | | Intent(意图识别) | function_call(函数调用) | function_call(函数调用) | | Memory(记忆功能) | mem_local_short(本地短期记忆) | mem_local_short(本地短期记忆) | @@ -260,7 +260,7 @@ Websocket接口地址: wss://2662r3426b.vicp.fun/xiaozhi/v1/ --- ## 产品生态 👬 -小智是一个生态,当你使用这个产品时,也可以看看其他在这个生态圈的[优秀项目](https://github.com/78/xiaozhi-esp32?tab=readme-ov-file#%E7%9B%B8%E5%85%B3%E5%BC%80%E6%BA%90%E9%A1%B9%E7%9B%AE) +小智是一个生态,当你使用这个产品时,也可以看看其他在这个生态圈的[优秀项目](https://github.com/78/xiaozhi-esp32/blob/main/README_zh.md#%E7%9B%B8%E5%85%B3%E5%BC%80%E6%BA%90%E9%A1%B9%E7%9B%AE) --- diff --git a/README_de.md b/README_de.md index c3578fec..0e4c74d5 100644 --- a/README_de.md +++ b/README_de.md @@ -181,8 +181,8 @@ Dieses Projekt bietet zwei Bereitstellungsmethoden. Bitte wählen Sie basierend #### 🚀 Auswahl der Bereitstellungsmethode | Bereitstellungsmethode | Funktionen | Anwendungsszenarien | Deployment-Dokumente | Konfigurationsanforderungen | Video-Tutorials | |---------|------|---------|---------|---------|---------| -| **Vereinfachte Installation** | Intelligenter Dialog, IOT, MCP, visuelle Wahrnehmung | Umgebungen mit geringer Konfiguration, Daten in Konfigurationsdateien gespeichert, keine Datenbank erforderlich | [①Docker-Version](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E5%8F%AA%E8%BF%90%E8%A1%8Cserver) / [②Quellcode-Deployment](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E5%8F%AA%E8%BF%90%E8%A1%8Cserver)| 2 Kerne 4GB bei Verwendung von `FunASR`, 2 Kerne 2GB bei allen APIs | - | -| **Vollständige Modulinstallation** | Intelligenter Dialog, IOT, MCP-Endpunkte, Stimmabdruckerkennung, visuelle Wahrnehmung, OTA, intelligente Steuerkonsole | Vollständige Funktionserfahrung, Daten in Datenbank gespeichert |[①Docker-Version](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) / [②Quellcode-Deployment](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) / [③Quellcode-Deployment Auto-Update-Tutorial](./docs/dev-ops-integration.md) | 4 Kerne 8GB bei Verwendung von `FunASR`, 2 Kerne 4GB bei allen APIs| [Video-Tutorial für lokalen Quellcode-Start](https://www.bilibili.com/video/BV1wBJhz4Ewe) | +| **Vereinfachte Installation** | Intelligenter Dialog, Einzel-Agenten-Verwaltung | Umgebungen mit geringer Konfiguration, Daten in Konfigurationsdateien gespeichert, keine Datenbank erforderlich | [①Docker-Version](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E5%8F%AA%E8%BF%90%E8%A1%8Cserver) / [②Quellcode-Deployment](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E5%8F%AA%E8%BF%90%E8%A1%8Cserver)| 2 Kerne 4GB bei Verwendung von `FunASR`, 2 Kerne 2GB bei allen APIs | - | +| **Vollständige Modulinstallation** | Intelligenter Dialog, Mehrbenutzerverwaltung, Mehr-Agenten-Verwaltung, Intelligente Steuerkonsole-Bedienung | Vollständige Funktionserfahrung, Daten in Datenbank gespeichert |[①Docker-Version](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) / [②Quellcode-Deployment](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) / [③Quellcode-Deployment Auto-Update-Tutorial](./docs/dev-ops-integration.md) | 4 Kerne 8GB bei Verwendung von `FunASR`, 2 Kerne 4GB bei allen APIs| [Video-Tutorial für lokalen Quellcode-Start](https://www.bilibili.com/video/BV1wBJhz4Ewe) | Häufige Fragen und entsprechende Tutorials finden Sie unter [diesem Link](./docs/FAQ.md) @@ -209,10 +209,10 @@ Websocket-Schnittstellenadresse: wss://2662r3426b.vicp.fun/xiaozhi/v1/ | Modulname | Einstiegslevel Kostenlose Einstellungen | Streaming-Konfiguration | |:---:|:---:|:---:| -| ASR (Spracherkennung) | FunASR (Lokal) | 👍FunASR (Lokaler GPU-Modus) | -| LLM (Großes Modell) | ChatGLMLLM (Zhipu glm-4-flash) | 👍AliLLM (qwen3-235b-a22b-instruct-2507) oder 👍DoubaoLLM (doubao-1-5-pro-32k-250115) | -| VLLM (Vision Large Model) | ChatGLMVLLM (Zhipu glm-4v-flash) | 👍QwenVLVLLM (Qwen qwen2.5-vl-3b-instructh) | -| TTS (Sprachsynthese) | ✅LinkeraiTTS (Lingxi-Streaming) | 👍HuoshanDoubleStreamTTS (Volcano Dual-Stream-Sprachsynthese) oder 👍AliyunStreamTTS (Alibaba Cloud Streaming-Sprachsynthese) | +| ASR (Spracherkennung) | FunASR (Lokal) | 👍XunfeiStreamASR (Xunfei-Streaming) | +| LLM (Großes Modell) | glm-4-flash (Zhipu) | 👍qwen-flash (Alibaba Bailian) | +| VLLM (Vision Large Model) | glm-4v-flash (Zhipu) | 👍qwen2.5-vl-3b-instructh (Alibaba Bailian) | +| TTS (Sprachsynthese) | ✅LinkeraiTTS (Lingxi-Streaming) | 👍HuoshanDoubleStreamTTS (Volcano-Streaming) | | Intent (Absichtserkennung) | function_call (Funktionsaufruf) | function_call (Funktionsaufruf) | | Memory (Gedächtnisfunktion) | mem_local_short (Lokales Kurzzeitgedächtnis) | mem_local_short (Lokales Kurzzeitgedächtnis) | @@ -258,7 +258,7 @@ Wenn Sie ein Softwareentwickler sind, finden Sie hier einen [Offenen Brief an En --- ## Produktökosystem 👬 -Xiaozhi ist ein Ökosystem. Wenn Sie dieses Produkt verwenden, können Sie sich auch andere [hervorragende Projekte](https://github.com/78/xiaozhi-esp32?tab=readme-ov-file#%E7%9B%B8%E5%85%B3%E5%BC%80%E6%BA%90%E9%A1%B9%E7%9B%AE) in diesem Ökosystem ansehen +Xiaozhi ist ein Ökosystem. Wenn Sie dieses Produkt verwenden, können Sie sich auch andere [hervorragende Projekte](https://github.com/78/xiaozhi-esp32?tab=readme-ov-file#related-open-source-projects) in diesem Ökosystem ansehen --- diff --git a/README_en.md b/README_en.md index e8437832..196cf16f 100644 --- a/README_en.md +++ b/README_en.md @@ -181,8 +181,8 @@ This project provides two deployment methods. Please choose based on your specif #### 🚀 Deployment Method Selection | Deployment Method | Features | Applicable Scenarios | Deployment Docs | Configuration Requirements | Video Tutorials | |---------|------|---------|---------|---------|---------| -| **Simplified Installation** | Intelligent dialogue, IOT, MCP, visual perception | Low-configuration environments, data stored in config files, no database required | [①Docker Version](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E5%8F%AA%E8%BF%90%E8%A1%8Cserver) / [②Source Code Deployment](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E5%8F%AA%E8%BF%90%E8%A1%8Cserver)| 2 cores 4GB if using `FunASR`, 2 cores 2GB if all APIs | - | -| **Full Module Installation** | Intelligent dialogue, IOT, MCP endpoints, voiceprint recognition, visual perception, OTA, intelligent control console | Complete functionality experience, data stored in database |[①Docker Version](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) / [②Source Code Deployment](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) / [③Source Code Deployment Auto-Update Tutorial](./docs/dev-ops-integration.md) | 4 cores 8GB if using `FunASR`, 2 cores 4GB if all APIs| [Local Source Code Startup Video Tutorial](https://www.bilibili.com/video/BV1wBJhz4Ewe) | +| **Simplified Installation** | Intelligent dialogue, single agent management | Low-configuration environments, data stored in config files, no database required | [①Docker Version](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E5%8F%AA%E8%BF%90%E8%A1%8Cserver) / [②Source Code Deployment](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E5%8F%AA%E8%BF%90%E8%A1%8Cserver)| 2 cores 4GB if using `FunASR`, 2 cores 2GB if all APIs | - | +| **Full Module Installation** | Intelligent dialogue, multi-user management, multi-agent management, intelligent console interface operation | Complete functionality experience, data stored in database |[①Docker Version](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) / [②Source Code Deployment](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) / [③Source Code Deployment Auto-Update Tutorial](./docs/dev-ops-integration.md) | 4 cores 8GB if using `FunASR`, 2 cores 4GB if all APIs| [Local Source Code Startup Video Tutorial](https://www.bilibili.com/video/BV1wBJhz4Ewe) | > 💡 Note: Below is a test platform deployed with the latest code. You can burn and test if needed. Concurrent users: 6, data will be cleared daily. @@ -208,10 +208,10 @@ Websocket Interface Address: wss://2662r3426b.vicp.fun/xiaozhi/v1/ | Module Name | Entry Level Free Settings | Streaming Configuration | |:---:|:---:|:---:| -| ASR(Speech Recognition) | FunASR(Local) | 👍FunASRServer or 👍DoubaoStreamASR | -| LLM(Large Model) | ChatGLMLLM(Zhipu glm-4-flash) | 👍DoubaoLLM(Volcano doubao-1-5-pro-32k-250115) | -| VLLM(Vision Large Model) | ChatGLMVLLM(Zhipu glm-4v-flash) | 👍QwenVLVLLM(Qwen qwen2.5-vl-3b-instructh) | -| TTS(Speech Synthesis) | ✅LinkeraiTTS(Lingxi streaming) | 👍HuoshanDoubleStreamTTS(Volcano dual-stream speech synthesis) | +| ASR(Speech Recognition) | FunASR(Local) | 👍XunfeiStreamASR(Xunfei Streaming) | +| LLM(Large Model) | glm-4-flash(Zhipu) | 👍qwen-flash(Alibaba Bailian) | +| VLLM(Vision Large Model) | glm-4v-flash(Zhipu) | 👍qwen2.5-vl-3b-instructh(Alibaba Bailian) | +| TTS(Speech Synthesis) | ✅LinkeraiTTS(Lingxi streaming) | 👍HuoshanDoubleStreamTTS(Volcano Streaming) | | Intent(Intent Recognition) | function_call(Function calling) | function_call(Function calling) | | Memory(Memory function) | mem_local_short(Local short-term memory) | mem_local_short(Local short-term memory) | @@ -256,7 +256,7 @@ If you are a software developer, here is an [Open Letter to Developers](docs/con --- ## Product Ecosystem 👬 -Xiaozhi is an ecosystem. When using this product, you can also check out other [excellent projects](https://github.com/78/xiaozhi-esp32?tab=readme-ov-file#%E7%9B%B8%E5%85%B3%E5%BC%80%E6%BA%90%E9%A1%B9%E7%9B%AE) in this ecosystem +Xiaozhi is an ecosystem. When using this product, you can also check out other [excellent projects](https://github.com/78/xiaozhi-esp32?tab=readme-ov-file#related-open-source-projects) in this ecosystem | Project Name | Project Address | Project Description | |:---------------------|:--------|:--------| diff --git a/README_vi.md b/README_vi.md index 92697551..e6d5f489 100644 --- a/README_vi.md +++ b/README_vi.md @@ -182,8 +182,8 @@ Dự án này cung cấp hai phương pháp triển khai, vui lòng chọn theo #### 🚀 Lựa chọn phương pháp triển khai | Phương pháp triển khai | Đặc điểm | Tình huống áp dụng | Tài liệu triển khai | Yêu cầu cấu hình | Video hướng dẫn | |---------|------|---------|---------|---------|---------| -| **Cài đặt tối giản** | Đối thoại thông minh, IOT, MCP, cảm nhận thị giác | Môi trường cấu hình thấp, dữ liệu lưu trong tệp cấu hình, không cần cơ sở dữ liệu | [①Phiên bản Docker](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E5%8F%AA%E8%BF%90%E8%A1%8Cserver) / [②Triển khai mã nguồn](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E5%8F%AA%E8%BF%90%E8%A1%8Cserver)| 2 nhân 4GB nếu dùng `FunASR`, 2 nhân 2GB nếu toàn API | - | -| **Cài đặt toàn bộ module** | Đối thoại thông minh, IOT, điểm truy cập MCP, nhận dạng giọng nói, cảm nhận thị giác, OTA, bảng điều khiển thông minh | Trải nghiệm đầy đủ tính năng, dữ liệu lưu trong cơ sở dữ liệu |[①Phiên bản Docker](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) / [②Triển khai mã nguồn](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) / [③Hướng dẫn tự động cập nhật triển khai mã nguồn](./docs/dev-ops-integration.md) | 4 nhân 8GB nếu dùng `FunASR`, 2 nhân 4GB nếu toàn API| [Video hướng dẫn khởi động mã nguồn cục bộ](https://www.bilibili.com/video/BV1wBJhz4Ewe) | +| **Cài đặt tối giản** | Đối thoại thông minh, quản lý đơn tác nhân | Môi trường cấu hình thấp, dữ liệu lưu trong tệp cấu hình, không cần cơ sở dữ liệu | [①Phiên bản Docker](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E5%8F%AA%E8%BF%90%E8%A1%8Cserver) / [②Triển khai mã nguồn](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E5%8F%AA%E8%BF%90%E8%A1%8Cserver)| 2 nhân 4GB nếu dùng `FunASR`, 2 nhân 2GB nếu toàn API | - | +| **Cài đặt toàn bộ module** | Đối thoại thông minh, quản lý đa người dùng, quản lý đa tác nhân, bảng điều khiển thông minh | Trải nghiệm đầy đủ tính năng, dữ liệu lưu trong cơ sở dữ liệu |[①Phiên bản Docker](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) / [②Triển khai mã nguồn](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) / [③Hướng dẫn tự động cập nhật triển khai mã nguồn](./docs/dev-ops-integration.md) | 4 nhân 8GB nếu dùng `FunASR`, 2 nhân 4GB nếu toàn API| [Video hướng dẫn khởi động mã nguồn cục bộ](https://www.bilibili.com/video/BV1wBJhz4Ewe) | Câu hỏi thường gặp và hướng dẫn liên quan, vui lòng tham khảo [liên kết này](./docs/FAQ.md) @@ -210,10 +210,10 @@ Công cụ kiểm tra dịch vụ: https://2662r3426b.vicp.fun/test/ | Tên module | Cài đặt miễn phí cho người mới | Cấu hình streaming | |:---:|:---:|:---:| -| ASR(Nhận dạng giọng nói) | FunASR(Local) | 👍FunASR(Chế độ GPU cục bộ) | -| LLM(Mô hình lớn) | ChatGLMLLM(Zhipu glm-4-flash) | 👍AliLLM(qwen3-235b-a22b-instruct-2507) hoặc 👍DoubaoLLM(doubao-1-5-pro-32k-250115) | -| VLLM(Mô hình lớn thị giác) | ChatGLMVLLM(Zhipu glm-4v-flash) | 👍QwenVLVLLM(Qwen qwen2.5-vl-3b-instructh) | -| TTS(Tổng hợp giọng nói) | ✅LinkeraiTTS(Lingxi streaming) | 👍HuoshanDoubleStreamTTS(Tổng hợp giọng nói streaming kép Volcano) hoặc 👍AliyunStreamTTS(Tổng hợp giọng nói streaming Alibaba Cloud) | +| ASR(Nhận dạng giọng nói) | FunASR(Local) | 👍XunfeiStreamASR(Xunfei Streaming) | +| LLM(Mô hình lớn) | glm-4-flash(Zhipu) | 👍qwen-flash(Alibaba Bailian) | +| VLLM(Mô hình lớn thị giác) | glm-4v-flash(Zhipu) | 👍qwen2.5-vl-3b-instructh(Alibaba Bailian) | +| TTS(Tổng hợp giọng nói) | ✅LinkeraiTTS(Lingxi streaming) | 👍HuoshanDoubleStreamTTS(Volcano Streaming) | | Intent(Nhận dạng ý định) | function_call(Gọi hàm) | function_call(Gọi hàm) | | Memory(Chức năng bộ nhớ) | mem_local_short(Bộ nhớ ngắn hạn cục bộ) | mem_local_short(Bộ nhớ ngắn hạn cục bộ) | @@ -259,7 +259,7 @@ Nếu bạn là một nhà phát triển phần mềm, đây có một [Lá thư --- ## Hệ sinh thái sản phẩm 👬 -Xiaozhi là một hệ sinh thái, khi bạn sử dụng sản phẩm này, bạn cũng có thể xem các [dự án xuất sắc](https://github.com/78/xiaozhi-esp32?tab=readme-ov-file#%E7%9B%B8%E5%85%B3%E5%BC%80%E6%BA%90%E9%A1%B9%E7%9B%AE) khác trong hệ sinh thái này +Xiaozhi là một hệ sinh thái, khi bạn sử dụng sản phẩm này, bạn cũng có thể xem các [dự án xuất sắc](https://github.com/78/xiaozhi-esp32?tab=readme-ov-file#related-open-source-projects) khác trong hệ sinh thái này --- diff --git a/docs/FAQ.md b/docs/FAQ.md index 83df8c2c..69494c3d 100644 --- a/docs/FAQ.md +++ b/docs/FAQ.md @@ -38,10 +38,10 @@ conda install conda-forge::ffmpeg | 模块名称 | 入门全免费设置 | 流式配置 | |:---:|:---:|:---:| -| ASR(语音识别) | FunASR(本地) | 👍FunASR(本地GPU模式) | -| LLM(大模型) | ChatGLMLLM(智谱glm-4-flash) | 👍AliLLM(qwen3-235b-a22b-instruct-2507) 或 👍DoubaoLLM(doubao-1-5-pro-32k-250115) | -| VLLM(视觉大模型) | ChatGLMVLLM(智谱glm-4v-flash) | 👍QwenVLVLLM(千问qwen2.5-vl-3b-instructh) | -| TTS(语音合成) | ✅LinkeraiTTS(灵犀流式) | 👍HuoshanDoubleStreamTTS(火山双流式语音合成) 或 👍AliyunStreamTTS(阿里云流式语音合成) | +| ASR(语音识别) | FunASR(本地) | 👍XunfeiStreamASR(讯飞流式) | +| LLM(大模型) | glm-4-flash(智谱) | 👍qwen-flash(阿里百炼) | +| VLLM(视觉大模型) | glm-4v-flash(智谱) | 👍qwen2.5-vl-3b-instructh(阿里百炼) | +| TTS(语音合成) | ✅LinkeraiTTS(灵犀流式) | 👍HuoshanDoubleStreamTTS(火山流式) | | Intent(意图识别) | function_call(函数调用) | function_call(函数调用) | | Memory(记忆功能) | mem_local_short(本地短期记忆) | mem_local_short(本地短期记忆) | @@ -69,6 +69,7 @@ VAD: ### 9、编译固件相关教程 1、[如何自己编译小智固件](./firmware-build.md)
2、[如何基于虾哥编译好的固件修改OTA地址](./firmware-setting.md)
+3、[单模块部署如何配置固件OTA自动升级](./ota-upgrade-guide.md)
### 10、拓展相关教程 1、[如何开启手机号码注册智控台](./ali-sms-integration.md)
@@ -80,6 +81,7 @@ VAD: 7、[如何开启声纹识别](./voiceprint-integration.md)
8、[新闻插件源配置指南](./newsnow_plugin_config.md)
9、[知识库ragflow集成指南](./ragflow-integration.md)
+10、[如何部署上下文源](./context-provider-integration.md)
### 11、语音克隆、本地语音部署相关教程 1、[如何在智控台克隆音色](./huoshan-streamTTS-voice-cloning.md)
diff --git a/docs/context-provider-integration.md b/docs/context-provider-integration.md new file mode 100644 index 00000000..d8a9049c --- /dev/null +++ b/docs/context-provider-integration.md @@ -0,0 +1,224 @@ +# 上下文源使用教程 + +## 概述 + +`上下文源`,就是为小智系统提示词的上下文添加【数据源】。 + +`上下文源` 在小智在唤醒那一刻,获取外部系统的数据,并将其动态注入到大模型的系统提示词(System Prompt)中。 +让其做到唤醒时感知世界某个事物的状态。 + +它和MCP、记忆有本质的区别:`上下文源`是强制让小智感知世界的数据;`记忆(Mem)`是让他知道之前聊了什么内容;`MCP(functionc all)`是当需要调用某项能力/知识的时候使用调用。 + +通过这个功能,在小智唤醒的一刹那,“感知”到: +- 人体健康传感器状态(体温、血压、血氧状态等) +- 业务系统的实时数据(服务器负载、待办数据、股票信息等) +- 任何可以通过 HTTP API 获取的文本信息 + +**注意**:该功能只是方便小智在唤醒的时候感知事物的状态,而如果想要小智唤醒后实时获取事物的状态,建议在此功能上再结合MCP工具的调用。 + +## 工作原理 + +1. **配置源**:用户配置一个或多个 HTTP API 地址。 +2. **触发请求**:当系统构建 Prompt 时,如果发现模板中包含 `{{ dynamic_context }}` 占位符,会请求所有配置的 API。 +3. **自动注入**:系统会自动将 API 返回的数据格式化为 Markdown 列表,替换 `{{ dynamic_context }}` 占位符。 + +## 接口规范 + +为了让小智正确解析数据,您的 API 需要满足以下规范: + +- **请求方式**:`GET` +- **请求头**:系统会自动添加 `device-id` 字段到 Request Header。 +- **响应格式**:必须返回 JSON 格式,且包含 `code` 和 `data` 字段。 + +### 响应示例 + +**情况 1:返回键值对** +```json +{ + "code": 0, + "msg": "success", + "data": { + "客厅温度": "26℃", + "客厅湿度": "45%", + "大门状态": "已关闭" + } +} +``` +*注入效果:* +```markdown + +- **客厅温度:** 26℃ +- **客厅湿度:** 45% +- **大门状态:** 已关闭 + +``` + +**情况 2:返回列表** +```json +{ + "code": 0, + "data": [ + "您有10个待办事项", + "当前汽车的行驶速度是100km每小时" + ] +} +``` +*注入效果:* +```markdown + +- 您有10个待办事项 +- 当前汽车的行驶速度是100km每小时 + +``` + +## 配置指南 + +### 方式 1:智控台配置(全模块部署) + +1. 登录智控台,进入**角色配置**页面。 +2. 找到**上下文源**配置项(点击“编辑源”按钮)。 +3. 点击**添加**,输入您的 API 地址。 +4. 如果 API 需要鉴权,可以在**请求头**部分添加 `Authorization` 或其他 Header。 +5. 保存配置。 + +### 方式 2:配置文件配置(单模块部署) + +编辑 `xiaozhi-server/data/.config.yaml` 文件,添加 `context_providers` 配置段: + +```yaml +# 上下文源配置 +context_providers: + - url: "http://api.example.com/data" + headers: + Authorization: "Bearer your-token" + - url: "http://another-api.com/data" +``` + +## 启用功能 + +默认情况下,系统的提示词模板文件(`data/.agent-base-prompt.txt`)中已经预置了 `{{ dynamic_context }}` 占位符,您无需手动添加。 + +**示例:** + +```markdown + +【重要!以下信息已实时提供,无需调用工具查询,请直接使用:】 +- **设备ID:** {{device_id}} +- **当前时间:** {{current_time}} +... +{{ dynamic_context }} + +``` + +**注意**:如果您不需要使用此功能,可以选择**不配置任何上下文源**,也可以从提示词模板文件中**删除** `{{ dynamic_context }}` 占位符。 + +## 附录:Mock 测试服务示例 + +为了方便您测试和开发,我们提供了一个简单的 Python Mock Server 脚本。您可以运行此脚本在本地模拟 API 接口。 + +**mock_api_server.py** + +```python +import http.server +import socketserver +import json +from urllib.parse import urlparse, parse_qs + +# 设置端口号 +PORT = 8081 + +class MockRequestHandler(http.server.SimpleHTTPRequestHandler): + def do_GET(self): + # 解析路径和参数 + parsed_path = urlparse(self.path) + path = parsed_path.path + query = parse_qs(parsed_path.query) + + response_data = {} + status_code = 200 + + print(f"收到请求: {path}, 参数: {query}") + + # Case 1: 模拟健康数据 (返回字典 Dict) + # 路径参数风格: /health + # device_id 从 Header 获取 + if path == "/health": + device_id = self.headers.get("device-id", "unknown_device") + print(f"device_id: {device_id}") + response_data = { + "code": 0, + "msg": "success", + "data": { + "测试设备ID": device_id, + "心率": "80 bpm", + "血压": "120/80 mmHg", + "状态": "良好" + } + } + + # Case 2: 模拟新闻列表 (返回列表 List) + # 无参数: /news/list + elif path == "/news/list": + response_data = { + "code": 0, + "msg": "success", + "data": [ + "今日头条:Python 3.14 发布", + "科技新闻:AI 助手改变生活", + "本地新闻:明日有大雨,记得带伞" + ] + } + + # Case 3: 模拟天气简报 (返回字符串 String) + # 无参数: /weather/simple + elif path == "/weather/simple": + response_data = { + "code": 0, + "msg": "success", + "data": "今日晴转多云,气温 20-25 度,空气质量优,适合出行。" + } + + # Case 4: 模拟设备详情 (Query参数风格) + # 参数风格: /device/info + # device_id 从 Header 获取 + elif path == "/device/info": + device_id = self.headers.get("device-id", "unknown_device") + response_data = { + "code": 0, + "msg": "success", + "data": { + "查询方式": "Header参数", + "设备ID": device_id, + "电量": "85%", + "固件": "v2.0.1" + } + } + + # Case 5: 404 Not Found + else: + status_code = 404 + response_data = {"error": "接口不存在"} + + # 发送响应 + self.send_response(status_code) + self.send_header('Content-type', 'application/json; charset=utf-8') + self.end_headers() + self.wfile.write(json.dumps(response_data, ensure_ascii=False).encode('utf-8')) + +# 启动服务 +# 允许地址重用,防止快速重启报错 +socketserver.TCPServer.allow_reuse_address = True +with socketserver.TCPServer(("", PORT), MockRequestHandler) as httpd: + print(f"==================================================") + print(f"Mock API Server 已启动: http://localhost:{PORT}") + print(f"可用接口列表:") + print(f"1. [字典] http://localhost:{PORT}/health") + print(f"2. [列表] http://localhost:{PORT}/news/list") + print(f"3. [文本] http://localhost:{PORT}/weather/simple") + print(f"4. [参数] http://localhost:{PORT}/device/info") + print(f"==================================================") + try: + httpd.serve_forever() + except KeyboardInterrupt: + print("\n服务已停止") +``` diff --git a/docs/huoshan-streamTTS-voice-cloning.md b/docs/huoshan-streamTTS-voice-cloning.md index c2d261d9..304a8c90 100644 --- a/docs/huoshan-streamTTS-voice-cloning.md +++ b/docs/huoshan-streamTTS-voice-cloning.md @@ -23,6 +23,8 @@ ### 2.将音色资源ID分配给系统账号 +使用超级管理员账号登录智控台,点击顶部`参数字典`,在下拉菜单中,点击`系统功能配置`页面。在页面上勾选`音色克隆`,点击保存配置。即可在顶部菜单看到`音色克隆`按钮。 + 使用超级管理员账号登录智控台,点击顶部【音色克隆】、【音色资源】。 点击新增按钮,在【平台名称】选择“火山双流式语音合成”; diff --git a/docs/mcp-endpoint-enable.md b/docs/mcp-endpoint-enable.md index 6b64ccac..2d86f802 100644 --- a/docs/mcp-endpoint-enable.md +++ b/docs/mcp-endpoint-enable.md @@ -71,6 +71,7 @@ docker logs -f mcp-endpoint-server 请你保留好上面两个`接口地址`,下一步要用到。 # 2、全模块部署时,怎么配置MCP接入点 +首先,你要开启MCP接入点功能。在智控台,点击顶部`参数字典`,在下拉菜单中,点击`系统功能配置`页面。在页面上勾选`MCP接入点`,点击`保存配置`。在`角色配置`页面,点击`编辑功能`按钮,即可看到`mcp接入点`功能。 如果你是全模块部署,使用管理员账号,登录智控台,点击顶部`参数字典`,选择`参数管理`功能。 diff --git a/docs/mqtt-gateway-integration.md b/docs/mqtt-gateway-integration.md index b2f3068c..2b96255f 100644 --- a/docs/mqtt-gateway-integration.md +++ b/docs/mqtt-gateway-integration.md @@ -76,6 +76,7 @@ MQTT_PORT=1883 # MQTT服务器端口 UDP_PORT=8884 # UDP服务器端口 API_PORT=8007 # 管理API端口 MQTT_SIGNATURE_KEY=test # MQTT签名密钥 +SERVER_SECRET=Te1st12134 # 服务器密钥,请保持和智控台(server.secret)一致或者和xiaozhi-server里(server.auth_key)保持一致 ``` 请注意`PUBLIC_IP`配置,确保其与实际公网IP一致,如果有域名就填域名。 @@ -85,6 +86,13 @@ MQTT_SIGNATURE_KEY=test # MQTT签名密钥 - 注意不要用简单的密码,比如`123456`、`test`等。 - 注意不要用简单的密码,比如`123456`、`test`等。 +`SERVER_SECRET` 是用生成websocket连接的认证信息。 + +1、如果你是全模块部署,且你的智控台的参数管理里`server.auth.enabled`设置成了`true`,那么,`SERVER_SECRET`需要和智控台(`server.secret`)保持一致。 + +2、如果你是单模块部署,且你在配置文件里把`server.auth.enabled`设置成了`true`,那么,`SERVER_SECRET`需要和配置文件里(`server.auth_key`)保持一致。 + + 6. 启动MQTT网关 ``` # 启动服务 diff --git a/docs/ota-upgrade-guide.md b/docs/ota-upgrade-guide.md new file mode 100644 index 00000000..0df4a2ff --- /dev/null +++ b/docs/ota-upgrade-guide.md @@ -0,0 +1,142 @@ +# 单模块部署固件OTA自动升级配置指南 + +本教程将指导你如何在**单模块部署**场景下配置固件OTA自动升级功能,实现设备固件的自动更新。 + +如果你已经使用**全模块部署**,请忽略本教程。 + +## 功能介绍 + +在单模块部署中,xiaozhi-server内置了OTA固件管理功能,可以自动检测设备版本并下发升级固件。系统会根据设备型号和当前版本,自动匹配并推送最新的固件版本。 + +## 前提条件 + +- 你已经成功进行**单模块部署**并运行xiaozhi-server +- 设备能够正常连接到服务器 + +## 第一步 准备固件文件 + +### 1. 创建固件存放目录 + +固件文件需要放在`data/bin/`目录下。如果该目录不存在,请手动创建: + +```bash +mkdir -p data/bin +``` + +### 2. 固件文件命名规则 + +固件文件必须遵循以下命名格式: + +``` +{设备型号}_{版本号}.bin +``` + +**命名规则说明:** +- `设备型号`:设备的型号名称,例如 `lichuang-dev`、`bread-compact-wifi` 等 +- `版本号`:固件版本号,必须以数字开头,支持数字、字母、点号、下划线和短横线,例如 `1.6.6`、`2.0.0` 等 +- 文件扩展名必须是 `.bin` + +**命名示例:** +``` +bread-compact-wifi_1.6.6.bin +lichuang-dev_2.0.0.bin +``` + +### 3. 放置固件文件 + +将准备好的固件文件(.bin文件)复制到`data/bin/`目录下: + +重要的事情说三遍:升级的bin文件是`xiaozhi.bin`,不是全量固件文件`merged-binary.bin`! + +重要的事情说三遍:升级的bin文件是`xiaozhi.bin`,不是全量固件文件`merged-binary.bin`! + +重要的事情说三遍:升级的bin文件是`xiaozhi.bin`,不是全量固件文件`merged-binary.bin`! + +```bash +cp xiaozhi.bin data/bin/设备型号_版本号.bin +``` + +例如: +```bash +cp xiaozhi.bin data/bin/bread-compact-wifi_1.6.6.bin +``` + +## 第二步 配置公网访问地址(仅公网部署需要) + +**注意:此步骤仅适用于单模块公网部署的场景。** + +如果你的xiaozhi-server是公网部署(使用公网IP或域名),**必须**配置`server.vision_explain`参数,因为OTA固件下载地址会使用该配置的域名和端口。 + +如果你是局域网部署,可以跳过此步骤。 + +### 为什么要配置这个参数? + +在单模块部署中,系统生成固件下载地址时,会使用`vision_explain`配置的域名和端口作为基础地址。如果不配置或配置错误,设备将无法访问固件下载地址。 + +### 配置方法 + +打开`data/.config.yaml`文件,找到`server`配置段,设置`vision_explain`参数: + +```yaml +server: + vision_explain: http://你的域名或IP:端口号/mcp/vision/explain +``` + +**配置示例:** + +局域网部署(默认): +```yaml +server: + vision_explain: http://192.168.1.100:8003/mcp/vision/explain +``` + +公网域名部署: +```yaml +server: + vision_explain: http://yourdomain.com:8003/mcp/vision/explain +``` + +### 注意事项 + +- 域名或IP必须是设备能够访问的地址 +- 如果使用Docker部署,不能使用Docker内部地址(如127.0.0.1或localhost) +- 如果你使用了nginx反向代理,请填写对外的地址和端口号,不是本项目运行的端口号 + + +## 常见问题 + +### 1. 设备收不到固件更新 + +**可能原因和解决方法:** + +- 检查固件文件命名是否符合规则:`{型号}_{版本号}.bin` +- 检查固件文件是否正确放置在`data/bin/`目录 +- 检查设备型号是否与固件文件名中的型号匹配 +- 检查固件版本号是否高于设备当前版本 +- 查看服务器日志,确认OTA请求是否正常处理 + +### 2. 设备报告下载地址无法访问 + +**可能原因和解决方法:** + +- 检查`server.vision_explain`配置的域名或IP是否正确 +- 确认端口号配置正确(默认8003) +- 如果是公网部署,确保设备能够访问该公网地址 +- 如果是Docker部署,确保不是使用了内部地址(127.0.0.1) +- 检查防火墙是否开放了对应端口 +- 如果你使用了nginx反向代理,请填写对外的地址和端口号,不是本项目运行的端口号 + +### 3. 如何确认设备当前版本 + +查看OTA请求日志,日志中会显示设备上报的版本号: + +``` +[ota_handler] - 设备 AA:BB:CC:DD:EE:FF 固件已是最新: 1.6.6 +``` + +### 4. 固件文件放置后没有生效 + +系统有30秒的缓存时间(默认),可以: +- 等待30秒后再让设备发起OTA请求 +- 重启xiaozhi-server服务 +- 调整`firmware_cache_ttl`配置为更短的时间 diff --git a/docs/ragflow-integration.md b/docs/ragflow-integration.md index c2cd20f5..7d2b4eab 100644 --- a/docs/ragflow-integration.md +++ b/docs/ragflow-integration.md @@ -156,6 +156,14 @@ services: 编辑`ragflow/docker`文件夹下的`.env`文件,找到以下配置,逐个搜索,逐个修改!逐个搜索,逐个修改! +下面对于`.env`文件的修改,60%的人会忽略`MYSQL_USER`配置导致ragflow启动不成功,因此,需要强调三次: + +强调第一次:如果你的`.env`文件如果没有`MYSQL_USER`配置,请在配置文件增加这项! + +强调第二次:如果你的`.env`文件如果没有`MYSQL_USER`配置,请在配置文件增加这项! + +强调第三次:如果你的`.env`文件如果没有`MYSQL_USER`配置,请在配置文件增加这项! + ``` env # 端口设置 SVR_WEB_HTTP_PORT=8008 # HTTP端口 @@ -230,9 +238,11 @@ docker-compose -f docker-compose.yml up -d 在弹框中,点击"Create new Key"按钮,生成一个API Key。复制这个`API Key`,你稍后会用到。 # 第二步 配置到智控台 -确保你的智控台版本是`0.8.7`或以上。使用超级管理员账号登录到智控台。在顶部导航栏中,点击`模型配置`,在左侧导航栏中,点击`知识库`。 +确保你的智控台版本是`0.8.7`或以上。使用超级管理员账号登录到智控台。 -在列表中找到`RAG_RAGFlow`,点击`编辑`按钮。 +首先,你要先开启知识库功能。在顶部导航栏中,点击`参数字典`,在下拉菜单中,点击`系统功能配置`页面。在页面上勾选`知识库`,点击`保存配置`。即可在导航栏看到`知识库`功能。 + +在顶部导航栏中,点击`模型配置`,在左侧导航栏中,点击`知识库`。在列表中找到`RAG_RAGFlow`,点击`编辑`按钮。 在`服务地址`中,填写`http://你的ragflow服务的局域网IP:8008`,例如我的ragflow服务的局域网IP是`192.168.1.100`,那么我就填写`http://192.168.1.100:8008`。 diff --git a/docs/voiceprint-integration.md b/docs/voiceprint-integration.md index fbc12b1f..bd0731e8 100644 --- a/docs/voiceprint-integration.md +++ b/docs/voiceprint-integration.md @@ -164,6 +164,8 @@ http://192.168.1.25:8005/voiceprint/health?key=abcd # 2、全模块部署时,怎么配置声纹识别 ## 第一步 配置接口 +首先,你要开启声纹识别功能。在智控台,点击顶部`参数字典`,在下拉菜单中,点击`系统功能配置`页面。在页面上勾选`声纹识别`,点击`保存配置`。即可在新建智能体的卡片上看到`声纹识别`按钮。 + 如果你是全模块部署,使用管理员账号,登录智控台,点击顶部`参数字典`,选择`参数管理`功能。 然后搜索参数`server.voice_print`,此时,它的值应该是`null`值。 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 a7deaba1..c3f498d4 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 @@ -141,6 +141,11 @@ public interface Constant { */ String SERVER_MQTT_SECRET = "server.mqtt_signature_key"; + /** + * WebSocket认证开关 + */ + String SERVER_AUTH_ENABLED = "server.auth.enabled"; + /** * 无记忆 */ @@ -299,7 +304,7 @@ public interface Constant { /** * 版本号 */ - public static final String VERSION = "0.8.8"; + public static final String VERSION = "0.8.10"; /** * 无效固件URL 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 3a969713..281472e0 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 @@ -159,4 +159,11 @@ public class RedisKeys { public static String getKnowledgeBaseCacheKey(String datasetId) { return "knowledge:base:" + datasetId; } + + /** + * 获取临时注册设备标记key + */ + public static String getTmpRegisterMacKey(String deviceId) { + return "tmp_register_mac:" + deviceId; + } } diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/controller/AgentController.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/controller/AgentController.java index a4ad0863..0052c7cc 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/agent/controller/AgentController.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/controller/AgentController.java @@ -44,6 +44,8 @@ import xiaozhi.modules.agent.entity.AgentEntity; import xiaozhi.modules.agent.entity.AgentTemplateEntity; import xiaozhi.modules.agent.service.AgentChatAudioService; import xiaozhi.modules.agent.service.AgentChatHistoryService; +import xiaozhi.modules.agent.service.AgentChatSummaryService; +import xiaozhi.modules.agent.service.AgentContextProviderService; import xiaozhi.modules.agent.service.AgentPluginMappingService; import xiaozhi.modules.agent.service.AgentService; import xiaozhi.modules.agent.service.AgentTemplateService; @@ -64,6 +66,8 @@ public class AgentController { private final AgentChatHistoryService agentChatHistoryService; private final AgentChatAudioService agentChatAudioService; private final AgentPluginMappingService agentPluginMappingService; + private final AgentContextProviderService agentContextProviderService; + private final AgentChatSummaryService agentChatSummaryService; private final RedisUtils redisUtils; @GetMapping("/list") @@ -117,6 +121,27 @@ public class AgentController { return new Result<>(); } + @PostMapping("/chat-summary/{sessionId}/save") + @Operation(summary = "根据会话ID生成聊天记录总结并保存(异步执行)") + public Result generateAndSaveChatSummary(@PathVariable String sessionId) { + try { + // 异步执行总结生成任务,立即返回成功响应 + new Thread(() -> { + try { + agentChatSummaryService.generateAndSaveChatSummary(sessionId); + System.out.println("异步执行会话 " + sessionId + " 的聊天记录总结完成"); + } catch (Exception e) { + System.err.println("异步执行会话 " + sessionId + " 的聊天记录总结失败: " + e.getMessage()); + } + }).start(); + + // 立即返回成功响应,不等待总结生成完成 + return new Result().ok(null); + } catch (Exception e) { + return new Result().error("启动异步总结生成任务失败: " + e.getMessage()); + } + } + @PutMapping("/{id}") @Operation(summary = "更新智能体") @RequiresPermissions("sys:role:normal") @@ -135,6 +160,8 @@ public class AgentController { agentChatHistoryService.deleteByAgentId(id, true, true); // 删除关联的插件 agentPluginMappingService.deleteByAgentId(id); + // 删除关联的上下文源配置 + agentContextProviderService.deleteByAgentId(id); // 再删除智能体 agentService.deleteById(id); return new Result<>(); @@ -182,6 +209,7 @@ public class AgentController { List result = agentChatHistoryService.getChatHistoryBySessionId(id, sessionId); return new Result>().ok(result); } + @GetMapping("/{id}/chat-history/user") @Operation(summary = "获取智能体聊天记录(用户)") @RequiresPermissions("sys:role:normal") diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/dao/AgentContextProviderDao.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/dao/AgentContextProviderDao.java new file mode 100644 index 00000000..d46ad3ab --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/dao/AgentContextProviderDao.java @@ -0,0 +1,9 @@ +package xiaozhi.modules.agent.dao; + +import org.apache.ibatis.annotations.Mapper; +import xiaozhi.common.dao.BaseDao; +import xiaozhi.modules.agent.entity.AgentContextProviderEntity; + +@Mapper +public interface AgentContextProviderDao extends BaseDao { +} diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/dao/AiAgentChatHistoryDao.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/dao/AiAgentChatHistoryDao.java index b75312d2..7c9bf1eb 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/agent/dao/AiAgentChatHistoryDao.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/dao/AiAgentChatHistoryDao.java @@ -1,6 +1,9 @@ package xiaozhi.modules.agent.dao; +import java.util.List; + import org.apache.ibatis.annotations.Mapper; +import org.apache.ibatis.annotations.Param; import com.baomidou.mybatisplus.core.mapper.BaseMapper; @@ -15,12 +18,6 @@ import xiaozhi.modules.agent.entity.AgentChatHistoryEntity; */ @Mapper public interface AiAgentChatHistoryDao extends BaseMapper { - /** - * 根据智能体ID删除音频 - * - * @param agentId 智能体ID - */ - void deleteAudioByAgentId(String agentId); /** * 根据智能体ID删除聊天历史记录 @@ -35,4 +32,19 @@ public interface AiAgentChatHistoryDao extends BaseMapper getAudioIdsByAgentId(String agentId); + + /** + * 批量删除音频 + * + * @param audioIds 音频ID列表 + */ + void deleteAudioByIds(@Param("audioIds") List audioIds); } diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/dto/AgentChatSummaryDTO.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/dto/AgentChatSummaryDTO.java new file mode 100644 index 00000000..f6fbe660 --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/dto/AgentChatSummaryDTO.java @@ -0,0 +1,45 @@ +package xiaozhi.modules.agent.dto; + +import io.swagger.v3.oas.annotations.media.Schema; +import lombok.Data; + +/** + * 智能体聊天记录总结DTO + */ +@Data +@Schema(description = "智能体聊天记录总结对象") +public class AgentChatSummaryDTO { + + @Schema(description = "会话ID") + private String sessionId; + + @Schema(description = "智能体ID") + private String agentId; + + @Schema(description = "总结内容") + private String summary; + + @Schema(description = "总结状态") + private boolean success; + + @Schema(description = "错误信息") + private String errorMessage; + + public AgentChatSummaryDTO() { + this.success = true; + } + + public AgentChatSummaryDTO(String sessionId, String agentId, String summary) { + this.sessionId = sessionId; + this.agentId = agentId; + this.summary = summary; + this.success = true; + } + + public AgentChatSummaryDTO(String sessionId, String errorMessage) { + this.sessionId = sessionId; + this.errorMessage = errorMessage; + this.success = false; + } + +} \ No newline at end of file diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/dto/AgentUpdateDTO.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/dto/AgentUpdateDTO.java index 0e3d9bc3..ebb29fbd 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/agent/dto/AgentUpdateDTO.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/dto/AgentUpdateDTO.java @@ -69,6 +69,9 @@ public class AgentUpdateDTO implements Serializable { @Schema(description = "排序", example = "1", nullable = true) private Integer sort; + @Schema(description = "上下文源配置", nullable = true) + private List contextProviders; + @Data @Schema(description = "插件函数信息") public static class FunctionInfo implements Serializable { diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/dto/ContextProviderDTO.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/dto/ContextProviderDTO.java new file mode 100644 index 00000000..0b8edfd8 --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/dto/ContextProviderDTO.java @@ -0,0 +1,19 @@ +package xiaozhi.modules.agent.dto; + +import java.io.Serializable; +import java.util.Map; + +import io.swagger.v3.oas.annotations.media.Schema; +import lombok.Data; + +@Data +@Schema(description = "上下文源配置DTO") +public class ContextProviderDTO implements Serializable { + private static final long serialVersionUID = 1L; + + @Schema(description = "URL地址") + private String url; + + @Schema(description = "请求头") + private Map headers; +} diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/entity/AgentContextProviderEntity.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/entity/AgentContextProviderEntity.java new file mode 100644 index 00000000..937556ef --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/entity/AgentContextProviderEntity.java @@ -0,0 +1,43 @@ +package xiaozhi.modules.agent.entity; + +import java.util.Date; +import java.util.List; + +import com.baomidou.mybatisplus.annotation.IdType; +import com.baomidou.mybatisplus.annotation.TableField; +import com.baomidou.mybatisplus.annotation.TableId; +import com.baomidou.mybatisplus.annotation.TableName; +import com.baomidou.mybatisplus.extension.handlers.JacksonTypeHandler; + +import io.swagger.v3.oas.annotations.media.Schema; +import lombok.Data; +import xiaozhi.modules.agent.dto.ContextProviderDTO; + +@Data +@TableName(value = "ai_agent_context_provider", autoResultMap = true) +@Schema(description = "智能体上下文源配置") +public class AgentContextProviderEntity { + + @TableId(type = IdType.ASSIGN_UUID) + @Schema(description = "主键") + private String id; + + @Schema(description = "智能体ID") + private String agentId; + + @Schema(description = "上下文源配置") + @TableField(typeHandler = JacksonTypeHandler.class) + private List contextProviders; + + @Schema(description = "创建者") + private Long creator; + + @Schema(description = "创建时间") + private Date createdAt; + + @Schema(description = "更新者") + private Long updater; + + @Schema(description = "更新时间") + private Date updatedAt; +} diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/service/AgentChatSummaryService.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/AgentChatSummaryService.java new file mode 100644 index 00000000..418ee9a0 --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/AgentChatSummaryService.java @@ -0,0 +1,15 @@ +package xiaozhi.modules.agent.service; + +/** + * 智能体聊天记录总结服务接口 + */ +public interface AgentChatSummaryService { + + /** + * 根据会话ID生成聊天记录总结并保存到智能体记忆 + * + * @param sessionId 会话ID + * @return 保存结果 + */ + boolean generateAndSaveChatSummary(String sessionId); +} \ No newline at end of file diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/service/AgentContextProviderService.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/AgentContextProviderService.java new file mode 100644 index 00000000..da759971 --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/AgentContextProviderService.java @@ -0,0 +1,25 @@ +package xiaozhi.modules.agent.service; + +import xiaozhi.common.service.BaseService; +import xiaozhi.modules.agent.entity.AgentContextProviderEntity; + +public interface AgentContextProviderService extends BaseService { + /** + * 根据智能体ID获取上下文源配置 + * @param agentId 智能体ID + * @return 上下文源配置实体 + */ + AgentContextProviderEntity getByAgentId(String agentId); + + /** + * 保存或更新上下文源配置 + * @param entity 实体 + */ + void saveOrUpdateByAgentId(AgentContextProviderEntity entity); + + /** + * 根据智能体ID删除上下文源配置 + * @param agentId 智能体ID + */ + void deleteByAgentId(String agentId); +} diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/service/biz/impl/AgentChatHistoryBizServiceImpl.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/biz/impl/AgentChatHistoryBizServiceImpl.java index 36b59993..8831bb1a 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/agent/service/biz/impl/AgentChatHistoryBizServiceImpl.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/biz/impl/AgentChatHistoryBizServiceImpl.java @@ -17,6 +17,7 @@ import xiaozhi.modules.agent.entity.AgentChatHistoryEntity; import xiaozhi.modules.agent.entity.AgentEntity; import xiaozhi.modules.agent.service.AgentChatAudioService; import xiaozhi.modules.agent.service.AgentChatHistoryService; +import xiaozhi.modules.agent.service.AgentChatSummaryService; import xiaozhi.modules.agent.service.AgentService; import xiaozhi.modules.agent.service.biz.AgentChatHistoryBizService; import xiaozhi.modules.device.entity.DeviceEntity; @@ -36,6 +37,7 @@ public class AgentChatHistoryBizServiceImpl implements AgentChatHistoryBizServic private final AgentService agentService; private final AgentChatHistoryService agentChatHistoryService; private final AgentChatAudioService agentChatAudioService; + private final AgentChatSummaryService agentChatSummaryService; private final RedisUtils redisUtils; private final DeviceService deviceService; @@ -50,7 +52,8 @@ public class AgentChatHistoryBizServiceImpl implements AgentChatHistoryBizServic public Boolean report(AgentChatHistoryReportDTO report) { String macAddress = report.getMacAddress(); Byte chatType = report.getChatType(); - Long reportTimeMillis = null != report.getReportTime() ? report.getReportTime() * 1000 : System.currentTimeMillis(); + Long reportTimeMillis = null != report.getReportTime() ? report.getReportTime() * 1000 + : System.currentTimeMillis(); log.info("小智设备聊天上报请求: macAddress={}, type={} reportTime={}", macAddress, chatType, reportTimeMillis); // 根据设备MAC地址查询对应的默认智能体,判断是否需要上报 @@ -105,7 +108,8 @@ public class AgentChatHistoryBizServiceImpl implements AgentChatHistoryBizServic /** * 组装上报数据 */ - private void saveChatText(AgentChatHistoryReportDTO report, String agentId, String macAddress, String audioId, Long reportTime) { + private void saveChatText(AgentChatHistoryReportDTO report, String agentId, String macAddress, String audioId, + Long reportTime) { // 构建聊天记录实体 AgentChatHistoryEntity entity = AgentChatHistoryEntity.builder() .macAddress(macAddress) diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentChatHistoryServiceImpl.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentChatHistoryServiceImpl.java index 76da2fc8..63b820c9 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentChatHistoryServiceImpl.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentChatHistoryServiceImpl.java @@ -84,7 +84,16 @@ public class AgentChatHistoryServiceImpl extends ServiceImpl audioIds = baseMapper.getAudioIdsByAgentId(agentId); + if (audioIds != null && !audioIds.isEmpty()) { + int batchSize = 1000; // 每批删除1000条 + for (int i = 0; i < audioIds.size(); i += batchSize) { + int end = Math.min(i + batchSize, audioIds.size()); + List batch = audioIds.subList(i, end); + baseMapper.deleteAudioByIds(batch); + } + } } if (deleteAudio && !deleteText) { baseMapper.deleteAudioIdByAgentId(agentId); @@ -107,7 +116,7 @@ public class AgentChatHistoryServiceImpl extends ServiceImpl pageParam = new Page<>(0, 50); diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentChatSummaryServiceImpl.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentChatSummaryServiceImpl.java new file mode 100644 index 00000000..f367891f --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentChatSummaryServiceImpl.java @@ -0,0 +1,423 @@ +package xiaozhi.modules.agent.service.impl; + +import java.util.ArrayList; +import java.util.List; +import java.util.Map; +import java.util.regex.Matcher; +import java.util.regex.Pattern; + +import org.apache.commons.lang3.StringUtils; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; +import org.springframework.stereotype.Service; + +import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper; + +import lombok.RequiredArgsConstructor; +import xiaozhi.modules.agent.dto.AgentChatHistoryDTO; +import xiaozhi.modules.agent.dto.AgentChatSummaryDTO; +import xiaozhi.modules.agent.dto.AgentMemoryDTO; +import xiaozhi.modules.agent.dto.AgentUpdateDTO; +import xiaozhi.modules.agent.entity.AgentChatHistoryEntity; +import xiaozhi.modules.agent.service.AgentChatHistoryService; +import xiaozhi.modules.agent.service.AgentChatSummaryService; +import xiaozhi.modules.agent.service.AgentService; +import xiaozhi.modules.agent.vo.AgentInfoVO; +import xiaozhi.modules.device.entity.DeviceEntity; +import xiaozhi.modules.device.service.DeviceService; +import xiaozhi.modules.llm.service.LLMService; +import xiaozhi.modules.model.entity.ModelConfigEntity; +import xiaozhi.modules.model.service.ModelConfigService; + +/** + * 智能体聊天记录总结服务实现类 + * 实现Python端mem_local_short.py中的总结逻辑 + */ +@Service +@RequiredArgsConstructor +public class AgentChatSummaryServiceImpl implements AgentChatSummaryService { + + private static final Logger log = LoggerFactory.getLogger(AgentChatSummaryServiceImpl.class); + + private final AgentChatHistoryService agentChatHistoryService; + private final AgentService agentService; + private final DeviceService deviceService; + private final LLMService llmService; + private final ModelConfigService modelConfigService; + + // 总结规则常量 + private static final int MAX_SUMMARY_LENGTH = 1800; // 最大总结长度 + private static final Pattern JSON_PATTERN = Pattern.compile("\\{.*?\\}", Pattern.DOTALL); + private static final Pattern DEVICE_CONTROL_PATTERN = Pattern.compile("设备控制|设备操作|控制设备|设备状态", + Pattern.CASE_INSENSITIVE); + private static final Pattern WEATHER_PATTERN = Pattern.compile("天气|温度|湿度|降雨|气象", Pattern.CASE_INSENSITIVE); + private static final Pattern DATE_PATTERN = Pattern.compile("日期|时间|星期|月份|年份", Pattern.CASE_INSENSITIVE); + + private AgentChatSummaryDTO generateChatSummary(String sessionId) { + try { + System.out.println("开始生成会话 " + sessionId + " 的聊天记录总结"); + + // 1. 根据sessionId获取聊天记录 + List chatHistory = getChatHistoryBySessionId(sessionId); + if (chatHistory == null || chatHistory.isEmpty()) { + return new AgentChatSummaryDTO(sessionId, "未找到该会话的聊天记录"); + } + + // 2. 获取智能体信息 + String agentId = getAgentIdFromSession(sessionId, chatHistory); + if (StringUtils.isBlank(agentId)) { + return new AgentChatSummaryDTO(sessionId, "无法获取智能体信息"); + } + + // 3. 提取关键对话内容 + List meaningfulMessages = extractMeaningfulMessages(chatHistory); + if (meaningfulMessages.isEmpty()) { + return new AgentChatSummaryDTO(sessionId, "没有有效的对话内容可总结"); + } + + // 4. 生成总结(generateSummaryFromMessages方法已包含长度限制逻辑) + String summary = generateSummaryFromMessages(meaningfulMessages, agentId); + + System.out.println("成功生成会话 " + sessionId + " 的聊天记录总结,长度: " + summary.length() + " 字符"); + return new AgentChatSummaryDTO(sessionId, agentId, summary); + + } catch (Exception e) { + System.err.println("生成会话 " + sessionId + " 的聊天记录总结时发生错误: " + e.getMessage()); + return new AgentChatSummaryDTO(sessionId, "生成总结时发生错误: " + e.getMessage()); + } + } + + @Override + public boolean generateAndSaveChatSummary(String sessionId) { + try { + // 1. 生成总结 + AgentChatSummaryDTO summaryDTO = generateChatSummary(sessionId); + if (!summaryDTO.isSuccess()) { + System.err.println("生成总结失败: " + summaryDTO.getErrorMessage()); + return false; + } + + // 2. 获取设备信息(通过会话关联的设备) + DeviceEntity device = getDeviceBySessionId(sessionId); + if (device == null) { + System.err.println("未找到与会话 " + sessionId + " 关联的设备"); + return false; + } + + // 3. 更新智能体记忆 + AgentMemoryDTO memoryDTO = new AgentMemoryDTO(); + memoryDTO.setSummaryMemory(summaryDTO.getSummary()); + + // 调用现有接口更新记忆 + agentService.updateAgentById(device.getAgentId(), + new AgentUpdateDTO() { + { + setSummaryMemory(summaryDTO.getSummary()); + } + }); + + System.out.println("成功保存会话 " + sessionId + " 的聊天记录总结到智能体 " + device.getAgentId()); + return true; + + } catch (Exception e) { + System.err.println("保存会话 " + sessionId + " 的聊天记录总结时发生错误: " + e.getMessage()); + return false; + } + } + + /** + * 根据会话ID获取聊天记录 + */ + private List getChatHistoryBySessionId(String sessionId) { + try { + // 这里需要根据sessionId获取聊天记录 + // 由于现有接口需要agentId,我们需要先找到关联的agentId + String agentId = findAgentIdBySessionId(sessionId); + if (StringUtils.isBlank(agentId)) { + return null; + } + return agentChatHistoryService.getChatHistoryBySessionId(agentId, sessionId); + } catch (Exception e) { + System.err.println("获取会话 " + sessionId + " 的聊天记录失败: " + e.getMessage()); + return null; + } + } + + /** + * 根据会话ID查找关联的智能体ID + */ + private String findAgentIdBySessionId(String sessionId) { + try { + // 查询该会话的第一条记录获取agentId + QueryWrapper wrapper = new QueryWrapper<>(); + wrapper.select("agent_id") + .eq("session_id", sessionId) + .last("LIMIT 1"); + + AgentChatHistoryEntity entity = agentChatHistoryService.getOne(wrapper); + return entity != null ? entity.getAgentId() : null; + } catch (Exception e) { + System.err.println("根据会话ID " + sessionId + " 查找智能体ID失败: " + e.getMessage()); + return null; + } + } + + /** + * 从会话中获取智能体ID + */ + private String getAgentIdFromSession(String sessionId, List chatHistory) { + // 直接从数据库查询智能体ID + return findAgentIdBySessionId(sessionId); + } + + /** + * 提取有意义的对话内容(只提取用户消息,排除AI回复) + */ + private List extractMeaningfulMessages(List chatHistory) { + List meaningfulMessages = new ArrayList<>(); + + for (AgentChatHistoryDTO message : chatHistory) { + // 只处理用户消息(chatType = 1) + if (message.getChatType() != null && message.getChatType() == 1) { + String content = extractContentFromMessage(message); + if (isMeaningfulMessage(content)) { + meaningfulMessages.add(content); + } + } + } + + return meaningfulMessages; + } + + /** + * 从消息中提取内容(处理JSON格式) + */ + private String extractContentFromMessage(AgentChatHistoryDTO message) { + String content = message.getContent(); + if (StringUtils.isBlank(content)) { + return ""; + } + + // 处理JSON格式内容(与前端ChatHistoryDialog.vue逻辑一致) + Matcher matcher = JSON_PATTERN.matcher(content); + if (matcher.find()) { + String jsonContent = matcher.group(); + // 简化处理:提取JSON中的文本内容 + return extractTextFromJson(jsonContent); + } + + return content; + } + + /** + * 从JSON中提取文本内容 + */ + private String extractTextFromJson(String jsonContent) { + // 简化处理:提取"content"字段的值 + Pattern contentPattern = Pattern.compile("\"content\"\s*:\s*\"([^\"]*)\""); + Matcher matcher = contentPattern.matcher(jsonContent); + if (matcher.find()) { + return matcher.group(1); + } + return jsonContent; + } + + /** + * 判断是否为有意义的消息 + */ + private boolean isMeaningfulMessage(String content) { + if (StringUtils.isBlank(content)) { + return false; + } + + // 排除设备控制信息 + if (DEVICE_CONTROL_PATTERN.matcher(content).find()) { + return false; + } + + // 排除日期天气等无关内容 + if (WEATHER_PATTERN.matcher(content).find() || DATE_PATTERN.matcher(content).find()) { + return false; + } + + // 排除过短的消息 + return content.length() >= 5; + } + + /** + * 从消息生成总结 + */ + private String generateSummaryFromMessages(List messages, String agentId) { + if (messages.isEmpty()) { + return "本次对话内容较少,没有需要总结的重要信息。"; + } + + // 构建完整的对话内容 + StringBuilder conversation = new StringBuilder(); + for (int i = 0; i < messages.size(); i++) { + conversation.append("消息").append(i + 1).append(": ").append(messages.get(i)).append("\n"); + } + + try { + // 获取当前智能体的历史记忆 + String historyMemory = getCurrentAgentMemory(agentId); + + // 调用LLM服务进行智能总结,传递agentId以获取正确的模型配置 + String summary = callJavaLLMForSummaryWithHistory(conversation.toString(), historyMemory, agentId); + + // 应用总结规则:限制最大长度 + if (summary.length() > MAX_SUMMARY_LENGTH) { + summary = summary.substring(0, MAX_SUMMARY_LENGTH) + "..."; + } + + return summary; + } catch (Exception e) { + System.err.println("调用Java端LLM服务失败: " + e.getMessage()); + throw new RuntimeException("LLM服务不可用,无法生成聊天总结"); + } + } + + /** + * 获取当前智能体的历史记忆 + */ + private String getCurrentAgentMemory(String agentId) { + try { + if (StringUtils.isBlank(agentId)) { + return null; + } + + // 获取智能体信息 + AgentInfoVO agentInfo = agentService.getAgentById(agentId); + if (agentInfo == null) { + return null; + } + + // 返回智能体的当前总结记忆 + return agentInfo.getSummaryMemory(); + } catch (Exception e) { + System.err.println("获取智能体历史记忆失败,agentId: " + agentId + ", 错误: " + e.getMessage()); + return null; + } + } + + /** + * 调用Java端LLM服务进行智能总结(支持历史记忆合并) + */ + private String callJavaLLMForSummaryWithHistory(String conversation, String historyMemory, String agentId) { + try { + // 获取智能体配置,从中提取记忆总结的模型ID + String modelId = getMemorySummaryModelId(agentId); + + if (StringUtils.isBlank(modelId)) { + System.out.println("未找到记忆总结的LLM模型配置,使用默认LLM服务"); + return llmService.generateSummaryWithHistory(conversation, historyMemory, null, null); + } + + // 使用指定的模型ID调用LLM服务(支持历史记忆合并) + String summary = llmService.generateSummaryWithHistory(conversation, historyMemory, null, modelId); + + if (StringUtils.isNotBlank(summary) && !summary.equals("服务暂不可用") && !summary.equals("总结生成失败")) { + return summary; + } + + throw new RuntimeException("Java端LLM服务返回异常: " + summary); + + } catch (Exception e) { + System.err.println("调用Java端LLM服务异常,agentId: " + agentId + ", 错误: " + e.getMessage()); + throw e; + } + } + + /** + * 调用Java端LLM服务进行智能总结 + */ + private String callJavaLLMForSummary(String conversation, String agentId) { + try { + // 获取智能体配置,从中提取记忆总结的模型ID + String modelId = getMemorySummaryModelId(agentId); + + if (StringUtils.isBlank(modelId)) { + System.out.println("未找到记忆总结的LLM模型配置,使用默认LLM服务"); + return llmService.generateSummary(conversation); + } + + // 使用指定的模型ID调用LLM服务 + String summary = llmService.generateSummaryWithModel(conversation, modelId); + + if (StringUtils.isNotBlank(summary) && !summary.equals("服务暂不可用") && !summary.equals("总结生成失败")) { + return summary; + } + + throw new RuntimeException("Java端LLM服务返回异常: " + summary); + + } catch (Exception e) { + System.err.println("调用Java端LLM服务异常,agentId: " + agentId + ", 错误: " + e.getMessage()); + throw e; + } + } + + /** + * 获取记忆总结的LLM模型ID + */ + private String getMemorySummaryModelId(String agentId) { + try { + if (StringUtils.isBlank(agentId)) { + return null; + } + + // 获取智能体信息 + AgentInfoVO agentInfo = agentService.getAgentById(agentId); + if (agentInfo == null) { + return null; + } + + // 获取智能体的记忆模型ID + String memModelId = agentInfo.getMemModelId(); + if (StringUtils.isBlank(memModelId)) { + return null; + } + + // 获取记忆模型配置 + ModelConfigEntity memModelConfig = modelConfigService.getModelByIdFromCache(memModelId); + if (memModelConfig == null || memModelConfig.getConfigJson() == null) { + return null; + } + + // 从记忆模型配置中提取对应的LLM模型ID + Map configMap = memModelConfig.getConfigJson(); + String llmModelId = (String) configMap.get("llm"); + + if (StringUtils.isBlank(llmModelId)) { + // 如果记忆模型没有配置独立的LLM,则使用智能体的默认LLM模型 + return agentInfo.getLlmModelId(); + } + + return llmModelId; + } catch (Exception e) { + System.err.println("获取记忆总结LLM模型ID失败,agentId: " + agentId + ", 错误: " + e.getMessage()); + return null; + } + } + + /** + * 根据会话ID获取设备信息 + */ + private DeviceEntity getDeviceBySessionId(String sessionId) { + try { + // 查询该会话的第一条记录获取macAddress + QueryWrapper wrapper = new QueryWrapper<>(); + wrapper.select("mac_address") + .eq("session_id", sessionId) + .last("LIMIT 1"); + + AgentChatHistoryEntity entity = agentChatHistoryService.getOne(wrapper); + if (entity != null && StringUtils.isNotBlank(entity.getMacAddress())) { + return deviceService.getDeviceByMacAddress(entity.getMacAddress()); + } + return null; + } catch (Exception e) { + System.err.println("根据会话ID " + sessionId + " 查找设备信息失败: " + e.getMessage()); + return null; + } + } +} \ No newline at end of file diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentContextProviderServiceImpl.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentContextProviderServiceImpl.java new file mode 100644 index 00000000..b68ab8c9 --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentContextProviderServiceImpl.java @@ -0,0 +1,35 @@ +package xiaozhi.modules.agent.service.impl; + +import org.springframework.stereotype.Service; + +import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper; + +import xiaozhi.common.service.impl.BaseServiceImpl; +import xiaozhi.modules.agent.dao.AgentContextProviderDao; +import xiaozhi.modules.agent.entity.AgentContextProviderEntity; +import xiaozhi.modules.agent.service.AgentContextProviderService; + +@Service +public class AgentContextProviderServiceImpl extends BaseServiceImpl implements AgentContextProviderService { + + @Override + public AgentContextProviderEntity getByAgentId(String agentId) { + return baseDao.selectOne(new QueryWrapper().eq("agent_id", agentId)); + } + + @Override + public void saveOrUpdateByAgentId(AgentContextProviderEntity entity) { + AgentContextProviderEntity exist = getByAgentId(entity.getAgentId()); + if (exist != null) { + entity.setId(exist.getId()); + updateById(entity); + } else { + insert(entity); + } + } + + @Override + public void deleteByAgentId(String agentId) { + baseDao.delete(new QueryWrapper().eq("agent_id", agentId)); + } +} diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentServiceImpl.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentServiceImpl.java index cb550b3e..0adf0873 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentServiceImpl.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentServiceImpl.java @@ -32,10 +32,12 @@ import xiaozhi.modules.agent.dao.AgentDao; import xiaozhi.modules.agent.dto.AgentCreateDTO; import xiaozhi.modules.agent.dto.AgentDTO; import xiaozhi.modules.agent.dto.AgentUpdateDTO; +import xiaozhi.modules.agent.entity.AgentContextProviderEntity; import xiaozhi.modules.agent.entity.AgentEntity; import xiaozhi.modules.agent.entity.AgentPluginMapping; import xiaozhi.modules.agent.entity.AgentTemplateEntity; import xiaozhi.modules.agent.service.AgentChatHistoryService; +import xiaozhi.modules.agent.service.AgentContextProviderService; import xiaozhi.modules.agent.service.AgentPluginMappingService; import xiaozhi.modules.agent.service.AgentService; import xiaozhi.modules.agent.service.AgentTemplateService; @@ -62,6 +64,7 @@ public class AgentServiceImpl extends BaseServiceImpl imp private final AgentChatHistoryService agentChatHistoryService; private final AgentTemplateService agentTemplateService; private final ModelProviderService modelProviderService; + private final AgentContextProviderService agentContextProviderService; @Override public PageData adminAgentList(Map params) { @@ -85,6 +88,13 @@ public class AgentServiceImpl extends BaseServiceImpl imp agent.setChatHistoryConf(Constant.ChatHistoryConfEnum.RECORD_TEXT_AUDIO.getCode()); } } + + // 查询上下文源配置 + AgentContextProviderEntity contextProviderEntity = agentContextProviderService.getByAgentId(id); + if (contextProviderEntity != null) { + agent.setContextProviders(contextProviderEntity.getContextProviders()); + } + // 无需额外查询插件列表,已通过SQL查询出来 return agent; } @@ -331,6 +341,14 @@ public class AgentServiceImpl extends BaseServiceImpl imp agentChatHistoryService.deleteByAgentId(existingEntity.getId(), true, false); } + // 更新上下文源配置 + if (dto.getContextProviders() != null) { + AgentContextProviderEntity contextEntity = new AgentContextProviderEntity(); + contextEntity.setAgentId(agentId); + contextEntity.setContextProviders(dto.getContextProviders()); + agentContextProviderService.saveOrUpdateByAgentId(contextEntity); + } + boolean b = validateLLMIntentParams(dto.getLlmModelId(), dto.getIntentModelId()); if (!b) { throw new RenException(ErrorCode.LLM_INTENT_PARAMS_MISMATCH); diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/vo/AgentInfoVO.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/vo/AgentInfoVO.java index d56c8bb4..3da8b71c 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/agent/vo/AgentInfoVO.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/vo/AgentInfoVO.java @@ -5,6 +5,7 @@ import com.baomidou.mybatisplus.extension.handlers.JacksonTypeHandler; import io.swagger.v3.oas.annotations.media.Schema; import lombok.Data; import lombok.EqualsAndHashCode; +import xiaozhi.modules.agent.dto.ContextProviderDTO; import xiaozhi.modules.agent.entity.AgentEntity; import xiaozhi.modules.agent.entity.AgentPluginMapping; @@ -21,4 +22,7 @@ public class AgentInfoVO extends AgentEntity @Schema(description = "插件列表Id") @TableField(typeHandler = JacksonTypeHandler.class) private List functions; + + @Schema(description = "上下文源配置") + private List contextProviders; } diff --git a/main/manager-api/src/main/java/xiaozhi/modules/config/service/impl/ConfigServiceImpl.java b/main/manager-api/src/main/java/xiaozhi/modules/config/service/impl/ConfigServiceImpl.java index 3ab3d977..ead30b70 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/config/service/impl/ConfigServiceImpl.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/config/service/impl/ConfigServiceImpl.java @@ -20,10 +20,12 @@ import xiaozhi.common.redis.RedisUtils; import xiaozhi.common.utils.ConvertUtils; import xiaozhi.common.utils.JsonUtils; import xiaozhi.modules.agent.dao.AgentVoicePrintDao; +import xiaozhi.modules.agent.entity.AgentContextProviderEntity; import xiaozhi.modules.agent.entity.AgentEntity; import xiaozhi.modules.agent.entity.AgentPluginMapping; import xiaozhi.modules.agent.entity.AgentTemplateEntity; import xiaozhi.modules.agent.entity.AgentVoicePrintEntity; +import xiaozhi.modules.agent.service.AgentContextProviderService; import xiaozhi.modules.agent.service.AgentMcpAccessPointService; import xiaozhi.modules.agent.service.AgentPluginMappingService; import xiaozhi.modules.agent.service.AgentService; @@ -53,6 +55,7 @@ public class ConfigServiceImpl implements ConfigService { private final TimbreService timbreService; private final AgentPluginMappingService agentPluginMappingService; private final AgentMcpAccessPointService agentMcpAccessPointService; + private final AgentContextProviderService agentContextProviderService; private final VoiceCloneService cloneVoiceService; private final AgentVoicePrintDao agentVoicePrintDao; @@ -103,6 +106,15 @@ public class ConfigServiceImpl implements ConfigService { @Override public Map getAgentModels(String macAddress, Map selectedModule) { + // 检查是否为管理控制台请求 + String redisKey = RedisKeys.getTmpRegisterMacKey(macAddress); + Object isAdminRequest = redisUtils.get(redisKey); + + if (isAdminRequest != null && "true".equals(isAdminRequest)) { + // 管理控制台请求,返回getConfig的结果 + redisUtils.delete(redisKey); // 使用后清理 + return (Map) getConfig(true); + } // 根据MAC地址查找设备 DeviceEntity device = deviceService.getDeviceByMacAddress(macAddress); if (device == null) { @@ -178,6 +190,13 @@ public class ConfigServiceImpl implements ConfigService { mcpEndpoint = mcpEndpoint.replace("/mcp/", "/call/"); result.put("mcp_endpoint", mcpEndpoint); } + + // 获取上下文源配置 + AgentContextProviderEntity contextProviderEntity = agentContextProviderService.getByAgentId(agent.getId()); + if (contextProviderEntity != null && contextProviderEntity.getContextProviders() != null && !contextProviderEntity.getContextProviders().isEmpty()) { + result.put("context_providers", contextProviderEntity.getContextProviders()); + } + // 获取声纹信息 buildVoiceprintConfig(agent.getId(), result); 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 index ff6069bf..78a17b02 100644 --- 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 @@ -72,10 +72,12 @@ public class DeviceController { return new Result().error(ErrorCode.MCA_NOT_NULL); } // 生成六位验证码 - String code = String.valueOf(Math.random()).substring(2, 8); - String key = RedisKeys.getDeviceCaptchaKey(code); + String code; + String key; String existsMac = null; do { + code = String.valueOf(Math.random()).substring(2, 8); + key = RedisKeys.getDeviceCaptchaKey(code); existsMac = (String) redisUtils.get(key); } while (StringUtils.isNotBlank(existsMac)); 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 index 3392c96f..5b56161f 100644 --- 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 @@ -98,4 +98,14 @@ public interface DeviceService extends BaseService { */ void updateDeviceConnectionInfo(String agentId, String deviceId, String appVersion); + /** + * 生成WebSocket认证token + * + * @param clientId 客户端ID + * @param username 用户名(通常为deviceId) + * @return 认证token字符串 + * @throws Exception 生成token时的异常 + */ + String generateWebSocketToken(String clientId, String username) throws Exception; + } \ 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 index 5c3f854d..5c807645 100644 --- 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 @@ -1,6 +1,8 @@ package xiaozhi.modules.device.service.impl; import java.nio.charset.StandardCharsets; +import java.security.InvalidKeyException; +import java.security.NoSuchAlgorithmException; import java.time.Instant; import java.util.Base64; import java.util.Date; @@ -169,7 +171,22 @@ public class DeviceServiceImpl extends BaseServiceImpl DeviceReportRespDTO.Websocket websocket = new DeviceReportRespDTO.Websocket(); // 从系统参数获取WebSocket URL,如果未配置则使用默认值 String wsUrl = sysParamsService.getValue(Constant.SERVER_WEBSOCKET, true); - websocket.setToken(""); + + // 检查是否启用认证并生成token + String authEnabled = sysParamsService.getValue(Constant.SERVER_AUTH_ENABLED, true); + if ("true".equalsIgnoreCase(authEnabled)) { + try { + // 生成token + String token = generateWebSocketToken(clientId, macAddress); + websocket.setToken(token); + } catch (Exception e) { + log.error("生成WebSocket token失败: {}", e.getMessage()); + websocket.setToken(""); + } + } else { + websocket.setToken(""); + } + if (StringUtils.isBlank(wsUrl) || wsUrl.equals("null")) { log.error("WebSocket地址未配置,请登录智控台,在参数管理找到【server.websocket】配置"); wsUrl = "ws://xiaozhi.server.com:8000/xiaozhi/v1/"; @@ -189,7 +206,7 @@ public class DeviceServiceImpl extends BaseServiceImpl // 添加MQTT UDP配置 // 从系统参数获取MQTT Gateway地址,仅在配置有效时使用 - String mqttUdpConfig = sysParamsService.getValue(Constant.SERVER_MQTT_GATEWAY, false); + String mqttUdpConfig = sysParamsService.getValue(Constant.SERVER_MQTT_GATEWAY, true); if (mqttUdpConfig != null && !mqttUdpConfig.equals("null") && !mqttUdpConfig.isEmpty()) { try { String groupId = deviceById != null && deviceById.getBoard() != null ? deviceById.getBoard() @@ -494,6 +511,40 @@ public class DeviceServiceImpl extends BaseServiceImpl return Base64.getEncoder().encodeToString(signature); } + /** + * 生成WebSocket认证token 遵循Python端AuthManager的实现逻辑:token = signature.timestamp + * + * @param clientId 客户端ID + * @param username 用户名 (通常为deviceId/macAddress) + * @return 认证token字符串 + */ + public String generateWebSocketToken(String clientId, String username) + throws NoSuchAlgorithmException, InvalidKeyException { + // 从系统参数获取密钥 + String secretKey = sysParamsService.getValue(Constant.SERVER_SECRET, false); + if (StringUtils.isBlank(secretKey)) { + throw new IllegalStateException("WebSocket认证密钥未配置(server.secret)"); + } + + // 获取当前时间戳(秒) + long timestamp = System.currentTimeMillis() / 1000; + + // 构建签名内容: clientId|username|timestamp + String content = String.format("%s|%s|%d", clientId, username, timestamp); + + // 生成HMAC-SHA256签名 + Mac hmac = Mac.getInstance("HmacSHA256"); + SecretKeySpec keySpec = new SecretKeySpec(secretKey.getBytes(StandardCharsets.UTF_8), "HmacSHA256"); + hmac.init(keySpec); + byte[] signature = hmac.doFinal(content.getBytes(StandardCharsets.UTF_8)); + + // Base64 URL-safe编码签名(去除填充符=) + String signatureBase64 = Base64.getUrlEncoder().withoutPadding().encodeToString(signature); + + // 返回格式: signature.timestamp + return String.format("%s.%d", signatureBase64, timestamp); + } + /** * 构建MQTT配置信息 * @@ -504,7 +555,7 @@ public class DeviceServiceImpl extends BaseServiceImpl private DeviceReportRespDTO.MQTT buildMqttConfig(String macAddress, String groupId) throws Exception { // 从环境变量或系统参数获取签名密钥 - String signatureKey = sysParamsService.getValue("server.mqtt_signature_key", false); + String signatureKey = sysParamsService.getValue("server.mqtt_signature_key", true); if (StringUtils.isBlank(signatureKey)) { log.warn("缺少MQTT_SIGNATURE_KEY,跳过MQTT配置生成"); return null; diff --git a/main/manager-api/src/main/java/xiaozhi/modules/llm/service/LLMService.java b/main/manager-api/src/main/java/xiaozhi/modules/llm/service/LLMService.java new file mode 100644 index 00000000..a3f10e9f --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/llm/service/LLMService.java @@ -0,0 +1,70 @@ +package xiaozhi.modules.llm.service; + +/** + * LLM服务接口 + * 支持多种大模型调用 + */ +public interface LLMService { + + /** + * 生成聊天记录总结 + * + * @param conversation 对话内容 + * @param promptTemplate 提示词模板 + * @return 总结结果 + */ + String generateSummary(String conversation, String promptTemplate); + + /** + * 生成聊天记录总结(使用默认提示词) + * + * @param conversation 对话内容 + * @return 总结结果 + */ + String generateSummary(String conversation); + + /** + * 生成聊天记录总结(指定模型ID) + * + * @param conversation 对话内容 + * @param modelId 模型ID + * @return 总结结果 + */ + String generateSummaryWithModel(String conversation, String modelId); + + /** + * 生成聊天记录总结(指定模型ID和提示词模板) + * + * @param conversation 对话内容 + * @param promptTemplate 提示词模板 + * @param modelId 模型ID + * @return 总结结果 + */ + String generateSummary(String conversation, String promptTemplate, String modelId); + + /** + * 生成聊天记录总结(包含历史记忆合并) + * + * @param conversation 对话内容 + * @param historyMemory 历史记忆 + * @param promptTemplate 提示词模板 + * @param modelId 模型ID + * @return 总结结果 + */ + String generateSummaryWithHistory(String conversation, String historyMemory, String promptTemplate, String modelId); + + /** + * 检查服务是否可用 + * + * @return 是否可用 + */ + boolean isAvailable(); + + /** + * 检查指定模型的服务是否可用 + * + * @param modelId 模型ID + * @return 是否可用 + */ + boolean isAvailable(String modelId); +} \ No newline at end of file diff --git a/main/manager-api/src/main/java/xiaozhi/modules/llm/service/impl/OpenAIStyleLLMServiceImpl.java b/main/manager-api/src/main/java/xiaozhi/modules/llm/service/impl/OpenAIStyleLLMServiceImpl.java new file mode 100644 index 00000000..55fda70d --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/llm/service/impl/OpenAIStyleLLMServiceImpl.java @@ -0,0 +1,305 @@ +package xiaozhi.modules.llm.service.impl; + +import java.util.HashMap; +import java.util.List; +import java.util.Map; + +import org.apache.commons.lang3.StringUtils; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.http.HttpEntity; +import org.springframework.http.HttpHeaders; +import org.springframework.http.HttpMethod; +import org.springframework.http.MediaType; +import org.springframework.http.ResponseEntity; +import org.springframework.stereotype.Service; +import org.springframework.web.client.RestTemplate; + +import cn.hutool.json.JSONArray; +import cn.hutool.json.JSONObject; +import cn.hutool.json.JSONUtil; +import lombok.extern.slf4j.Slf4j; +import xiaozhi.modules.llm.service.LLMService; +import xiaozhi.modules.model.entity.ModelConfigEntity; +import xiaozhi.modules.model.service.ModelConfigService; + +/** + * OpenAI风格API的LLM服务实现 + * 支持阿里云、DeepSeek、ChatGLM等兼容OpenAI API的模型 + */ +@Slf4j +@Service +public class OpenAIStyleLLMServiceImpl implements LLMService { + + @Autowired + private ModelConfigService modelConfigService; + + private final RestTemplate restTemplate = new RestTemplate(); + + private static final String DEFAULT_SUMMARY_PROMPT = "你是一个经验丰富的记忆总结者,擅长将对话内容进行总结摘要,遵循以下规则:\n1、总结用户的重要信息,以便在未来的对话中提供更个性化的服务\n2、不要重复总结,不要遗忘之前记忆,除非原来的记忆超过了1800字,否则不要遗忘、不要压缩用户的历史记忆\n3、用户操控的设备音量、播放音乐、天气、退出、不想对话等和用户本身无关的内容,这些信息不需要加入到总结中\n4、聊天内容中的今天的日期时间、今天的天气情况与用户事件无关的数据,这些信息如果当成记忆存储会影响后续对话,这些信息不需要加入到总结中\n5、不要把设备操控的成果结果和失败结果加入到总结中,也不要把用户的一些废话加入到总结中\n6、不要为了总结而总结,如果用户的聊天没有意义,请返回原来的历史记录也是可以的\n7、只需要返回总结摘要,严格控制在1800字内\n8、不要包含代码、xml,不需要解释、注释和说明,保存记忆时仅从对话提取信息,不要混入示例内容\n9、如果提供了历史记忆,请将新对话内容与历史记忆进行智能合并,保留有价值的历史信息,同时添加新的重要信息\n\n历史记忆:\n{history_memory}\n\n新对话内容:\n{conversation}"; + + @Override + public String generateSummary(String conversation) { + return generateSummary(conversation, null, null); + } + + @Override + public String generateSummaryWithModel(String conversation, String modelId) { + return generateSummary(conversation, null, modelId); + } + + @Override + public String generateSummary(String conversation, String promptTemplate, String modelId) { + if (!isAvailable()) { + log.warn("LLM服务不可用,无法生成总结"); + return "LLM服务不可用,无法生成总结"; + } + + try { + // 从智控台获取LLM模型配置 + ModelConfigEntity llmConfig; + if (modelId != null && !modelId.trim().isEmpty()) { + // 通过具体模型ID获取配置 + llmConfig = modelConfigService.getModelByIdFromCache(modelId); + } else { + // 保持向后兼容,使用默认配置 + llmConfig = getDefaultLLMConfig(); + } + + if (llmConfig == null || llmConfig.getConfigJson() == null) { + log.error("未找到可用的LLM模型配置,modelId: {}", modelId); + return "未找到可用的LLM模型配置"; + } + + JSONObject configJson = llmConfig.getConfigJson(); + String baseUrl = configJson.getStr("base_url"); + String model = configJson.getStr("model_name"); + String apiKey = configJson.getStr("api_key"); + Double temperature = configJson.getDouble("temperature"); + Integer maxTokens = configJson.getInt("max_tokens"); + + if (StringUtils.isBlank(baseUrl) || StringUtils.isBlank(apiKey)) { + log.error("LLM配置不完整,baseUrl或apiKey为空"); + return "LLM配置不完整,无法生成总结"; + } + + // 构建提示词 + String prompt = (promptTemplate != null ? promptTemplate : DEFAULT_SUMMARY_PROMPT).replace("{conversation}", + conversation); + + // 构建请求体 + Map requestBody = new HashMap<>(); + requestBody.put("model", model != null ? model : "gpt-3.5-turbo"); + + Map[] messages = new Map[1]; + Map message = new HashMap<>(); + message.put("role", "user"); + message.put("content", prompt); + messages[0] = message; + + requestBody.put("messages", messages); + requestBody.put("temperature", temperature != null ? temperature : 0.7); + requestBody.put("max_tokens", maxTokens != null ? maxTokens : 2000); + + // 发送HTTP请求 + HttpHeaders headers = new HttpHeaders(); + headers.setContentType(MediaType.APPLICATION_JSON); + headers.set("Authorization", "Bearer " + apiKey); + + HttpEntity> entity = new HttpEntity<>(requestBody, headers); + + // 构建完整的API URL + String apiUrl = baseUrl; + if (!apiUrl.endsWith("/chat/completions")) { + if (!apiUrl.endsWith("/")) { + apiUrl += "/"; + } + apiUrl += "chat/completions"; + } + + ResponseEntity response = restTemplate.exchange( + apiUrl, HttpMethod.POST, entity, String.class); + + if (response.getStatusCode().is2xxSuccessful()) { + JSONObject responseJson = JSONUtil.parseObj(response.getBody()); + JSONArray choices = responseJson.getJSONArray("choices"); + if (choices != null && choices.size() > 0) { + JSONObject choice = choices.getJSONObject(0); + JSONObject messageObj = choice.getJSONObject("message"); + return messageObj.getStr("content"); + } + } else { + log.error("LLM API调用失败,状态码:{},响应:{}", response.getStatusCode(), response.getBody()); + } + } catch (Exception e) { + log.error("调用LLM服务生成总结时发生异常,modelId: {}", modelId, e); + } + + return "生成总结失败,请稍后重试"; + } + + @Override + public String generateSummary(String conversation, String promptTemplate) { + return generateSummary(conversation, promptTemplate, null); + } + + @Override + public String generateSummaryWithHistory(String conversation, String historyMemory, String promptTemplate, + String modelId) { + if (!isAvailable()) { + log.warn("LLM服务不可用,无法生成总结"); + return "LLM服务不可用,无法生成总结"; + } + + try { + // 从智控台获取LLM模型配置 + ModelConfigEntity llmConfig; + if (modelId != null && !modelId.trim().isEmpty()) { + // 通过具体模型ID获取配置 + llmConfig = modelConfigService.getModelByIdFromCache(modelId); + } else { + // 保持向后兼容,使用默认配置 + llmConfig = getDefaultLLMConfig(); + } + + if (llmConfig == null || llmConfig.getConfigJson() == null) { + log.error("未找到可用的LLM模型配置,modelId: {}", modelId); + return "未找到可用的LLM模型配置"; + } + + JSONObject configJson = llmConfig.getConfigJson(); + String baseUrl = configJson.getStr("base_url"); + String model = configJson.getStr("model_name"); + String apiKey = configJson.getStr("api_key"); + + if (StringUtils.isBlank(baseUrl) || StringUtils.isBlank(apiKey)) { + log.error("LLM配置不完整,baseUrl或apiKey为空"); + return "LLM配置不完整,无法生成总结"; + } + + // 构建提示词,包含历史记忆 + String prompt = (promptTemplate != null ? promptTemplate : DEFAULT_SUMMARY_PROMPT) + .replace("{history_memory}", historyMemory != null ? historyMemory : "无历史记忆") + .replace("{conversation}", conversation); + + // 构建请求体 + Map requestBody = new HashMap<>(); + requestBody.put("model", model != null ? model : "gpt-3.5-turbo"); + + Map[] messages = new Map[1]; + Map message = new HashMap<>(); + message.put("role", "user"); + message.put("content", prompt); + messages[0] = message; + + requestBody.put("messages", messages); + requestBody.put("temperature", 0.2); + requestBody.put("max_tokens", 2000); + + // 发送HTTP请求 + HttpHeaders headers = new HttpHeaders(); + headers.setContentType(MediaType.APPLICATION_JSON); + headers.set("Authorization", "Bearer " + apiKey); + + HttpEntity> entity = new HttpEntity<>(requestBody, headers); + + // 构建完整的API URL + String apiUrl = baseUrl; + if (!apiUrl.endsWith("/chat/completions")) { + if (!apiUrl.endsWith("/")) { + apiUrl += "/"; + } + apiUrl += "chat/completions"; + } + + ResponseEntity response = restTemplate.exchange( + apiUrl, HttpMethod.POST, entity, String.class); + + if (response.getStatusCode().is2xxSuccessful()) { + JSONObject responseJson = JSONUtil.parseObj(response.getBody()); + JSONArray choices = responseJson.getJSONArray("choices"); + if (choices != null && choices.size() > 0) { + JSONObject choice = choices.getJSONObject(0); + JSONObject messageObj = choice.getJSONObject("message"); + return messageObj.getStr("content"); + } + } else { + log.error("LLM API调用失败,状态码:{},响应:{}", response.getStatusCode(), response.getBody()); + } + } catch (Exception e) { + log.error("调用LLM服务生成总结时发生异常,modelId: {}", modelId, e); + } + + return "生成总结失败,请稍后重试"; + } + + @Override + public boolean isAvailable() { + try { + ModelConfigEntity defaultLLMConfig = getDefaultLLMConfig(); + if (defaultLLMConfig == null || defaultLLMConfig.getConfigJson() == null) { + return false; + } + + JSONObject configJson = defaultLLMConfig.getConfigJson(); + String baseUrl = configJson.getStr("base_url"); + String apiKey = configJson.getStr("api_key"); + + return baseUrl != null && !baseUrl.trim().isEmpty() && + apiKey != null && !apiKey.trim().isEmpty(); + } catch (Exception e) { + log.error("检查LLM服务可用性时发生异常:", e); + return false; + } + } + + @Override + public boolean isAvailable(String modelId) { + try { + if (modelId == null || modelId.trim().isEmpty()) { + return isAvailable(); + } + + // 通过具体模型ID获取配置 + ModelConfigEntity modelConfig = modelConfigService.getModelByIdFromCache(modelId); + if (modelConfig == null || modelConfig.getConfigJson() == null) { + log.warn("未找到指定的LLM模型配置,modelId: {}", modelId); + return false; + } + + JSONObject configJson = modelConfig.getConfigJson(); + String baseUrl = configJson.getStr("base_url"); + String apiKey = configJson.getStr("api_key"); + + return baseUrl != null && !baseUrl.trim().isEmpty() && + apiKey != null && !apiKey.trim().isEmpty(); + } catch (Exception e) { + log.error("检查LLM服务可用性时发生异常,modelId: {}", modelId, e); + return false; + } + } + + /** + * 从智控台获取默认的LLM模型配置 + */ + private ModelConfigEntity getDefaultLLMConfig() { + try { + // 获取所有启用的LLM模型配置 + List llmConfigs = modelConfigService.getEnabledModelsByType("LLM"); + if (llmConfigs == null || llmConfigs.isEmpty()) { + return null; + } + + // 优先返回默认配置,如果没有默认配置则返回第一个启用的配置 + for (ModelConfigEntity config : llmConfigs) { + if (config.getIsDefault() != null && config.getIsDefault() == 1) { + return config; + } + } + + return llmConfigs.get(0); + } catch (Exception e) { + log.error("获取LLM模型配置时发生异常:", e); + return null; + } + } +} \ No newline at end of file diff --git a/main/manager-api/src/main/java/xiaozhi/modules/model/service/ModelConfigService.java b/main/manager-api/src/main/java/xiaozhi/modules/model/service/ModelConfigService.java index 95b64b20..9ae10761 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/model/service/ModelConfigService.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/model/service/ModelConfigService.java @@ -55,4 +55,12 @@ public interface ModelConfigService extends BaseService { * @return TTS平台列表(id和modelName) */ List> getTtsPlatformList(); + + /** + * 根据模型类型获取所有启用的模型配置 + * + * @param modelType 模型类型(如:LLM, TTS, ASR等) + * @return 启用的模型配置列表 + */ + List getEnabledModelsByType(String modelType); } diff --git a/main/manager-api/src/main/java/xiaozhi/modules/model/service/impl/ModelConfigServiceImpl.java b/main/manager-api/src/main/java/xiaozhi/modules/model/service/impl/ModelConfigServiceImpl.java index fbfc7723..2a291cf9 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/model/service/impl/ModelConfigServiceImpl.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/model/service/impl/ModelConfigServiceImpl.java @@ -502,4 +502,22 @@ public class ModelConfigServiceImpl extends BaseServiceImpl> getTtsPlatformList() { return modelConfigDao.getTtsPlatformList(); } + + /** + * 根据模型类型获取所有启用的模型配置 + */ + @Override + public List getEnabledModelsByType(String modelType) { + if (StringUtils.isBlank(modelType)) { + return null; + } + + List entities = modelConfigDao.selectList( + new QueryWrapper() + .eq("model_type", modelType) + .eq("is_enabled", 1) + .orderByAsc("sort")); + + return entities; + } } 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 051cbc29..29b9ad2f 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 @@ -89,7 +89,7 @@ public class ShiroConfig { filterMap.put("/config/**", "server"); filterMap.put("/agent/chat-history/report", "server"); filterMap.put("/agent/chat-history/download/**", "anon"); - filterMap.put("/agent/saveMemory/**", "server"); + filterMap.put("/agent/chat-summary/**", "server"); filterMap.put("/agent/play/**", "anon"); filterMap.put("/voiceClone/play/**", "anon"); filterMap.put("/**", "oauth2"); 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 a008c25b..1ab9809d 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 @@ -6,6 +6,7 @@ import java.util.HashMap; import java.util.List; import java.util.Map; +import org.apache.commons.lang3.StringUtils; import org.springframework.web.bind.annotation.GetMapping; import org.springframework.web.bind.annotation.PostMapping; import org.springframework.web.bind.annotation.PutMapping; @@ -23,7 +24,9 @@ import xiaozhi.common.exception.ErrorCode; import xiaozhi.common.exception.RenException; import xiaozhi.common.page.TokenDTO; import xiaozhi.common.user.UserDetail; +import xiaozhi.common.utils.JsonUtils; import xiaozhi.common.utils.Result; +import xiaozhi.common.utils.Sm2DecryptUtil; import xiaozhi.common.validator.AssertUtils; import xiaozhi.common.validator.ValidatorUtils; import xiaozhi.modules.security.dto.LoginDTO; @@ -32,8 +35,6 @@ 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.common.utils.Sm2DecryptUtil; -import org.apache.commons.lang3.StringUtils; import xiaozhi.modules.sys.dto.PasswordDTO; import xiaozhi.modules.sys.dto.RetrievePasswordDTO; import xiaozhi.modules.sys.dto.SysUserDTO; @@ -89,13 +90,13 @@ public class LoginController { @Operation(summary = "登录") public Result login(@RequestBody LoginDTO login) { String password = login.getPassword(); - + // 使用工具类解密并验证验证码 String actualPassword = Sm2DecryptUtil.decryptAndValidateCaptcha( password, login.getCaptchaId(), captchaService, sysParamsService); - + login.setPassword(actualPassword); - + // 按照用户名获取用户 SysUserDTO userDTO = sysUserService.getByUsername(login.getUsername()); // 判断用户是否存在 @@ -108,8 +109,6 @@ public class LoginController { } return sysUserTokenService.createToken(userDTO.getId()); } - - @PostMapping("/register") @Operation(summary = "注册") @@ -117,15 +116,15 @@ public class LoginController { if (!sysUserService.getAllowUserRegister()) { throw new RenException(ErrorCode.USER_REGISTER_DISABLED); } - + String password = login.getPassword(); - + // 使用工具类解密并验证验证码 String actualPassword = Sm2DecryptUtil.decryptAndValidateCaptcha( password, login.getCaptchaId(), captchaService, sysParamsService); - + login.setPassword(actualPassword); - + // 是否开启手机注册 Boolean isMobileRegister = sysParamsService .getValueObject(Constant.SysMSMParam.SERVER_ENABLE_MOBILE_REGISTER.getValue(), Boolean.class); @@ -204,11 +203,11 @@ public class LoginController { } String password = dto.getPassword(); - + // 使用工具类解密并验证验证码 String actualPassword = Sm2DecryptUtil.decryptAndValidateCaptcha( password, dto.getCaptchaId(), captchaService, sysParamsService); - + dto.setPassword(actualPassword); sysUserService.changePasswordDirectly(userDTO.getId(), dto.getPassword()); @@ -229,7 +228,7 @@ public class LoginController { config.put("beianIcpNum", sysParamsService.getValue(Constant.SysBaseParam.BEIAN_ICP_NUM.getValue(), true)); config.put("beianGaNum", sysParamsService.getValue(Constant.SysBaseParam.BEIAN_GA_NUM.getValue(), true)); config.put("name", sysParamsService.getValue(Constant.SysBaseParam.SERVER_NAME.getValue(), true)); - + // SM2公钥 String publicKey = sysParamsService.getValue(Constant.SM2_PUBLIC_KEY, true); if (StringUtils.isBlank(publicKey)) { @@ -237,6 +236,12 @@ public class LoginController { } config.put("sm2PublicKey", publicKey); + // 获取system-web.menu参数配置 + String menuConfig = sysParamsService.getValue("system-web.menu", true); + if (StringUtils.isNotBlank(menuConfig)) { + config.put("systemWebMenu", JsonUtils.parseObject(menuConfig, Object.class)); + } + return new Result>().ok(config); } } \ No newline at end of file diff --git a/main/manager-api/src/main/java/xiaozhi/modules/sys/controller/ServerSideManageController.java b/main/manager-api/src/main/java/xiaozhi/modules/sys/controller/ServerSideManageController.java index 893ea8e2..b467630d 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/sys/controller/ServerSideManageController.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/sys/controller/ServerSideManageController.java @@ -31,6 +31,8 @@ import xiaozhi.modules.sys.dto.ServerActionResponseDTO; import xiaozhi.modules.sys.enums.ServerActionEnum; import xiaozhi.modules.sys.service.SysParamsService; import xiaozhi.modules.sys.utils.WebSocketClientManager; +import xiaozhi.modules.device.service.DeviceService; +import xiaozhi.common.redis.RedisUtils; /** * 服务端管理控制器 @@ -41,6 +43,8 @@ import xiaozhi.modules.sys.utils.WebSocketClientManager; @AllArgsConstructor public class ServerSideManageController { private final SysParamsService sysParamsService; + private final DeviceService deviceService; + private final RedisUtils redisUtils; private static final ObjectMapper objectMapper; static { objectMapper = new ObjectMapper(); @@ -85,9 +89,22 @@ public class ServerSideManageController { return false; } String serverSK = sysParamsService.getValue(Constant.SERVER_SECRET, true); + + String deviceId = UUID.randomUUID().toString(); + String clientId = UUID.randomUUID().toString(); + + String redisKey = xiaozhi.common.redis.RedisKeys.getTmpRegisterMacKey(deviceId); + redisUtils.set(redisKey, "true", 300); // 5分钟有效期 + WebSocketHttpHeaders headers = new WebSocketHttpHeaders(); - headers.add("device-id", UUID.randomUUID().toString()); - headers.add("client-id", UUID.randomUUID().toString()); + headers.add("device-id", deviceId); + headers.add("client-id", clientId); + try { + String token = deviceService.generateWebSocketToken(clientId, deviceId); + headers.add("authorization", "Bearer " + token); + } catch (Exception e) { + throw new RenException(ErrorCode.WEB_SOCKET_CONNECT_FAILED); + } try (WebSocketClientManager client = new WebSocketClientManager.Builder() .connectTimeout(3, TimeUnit.SECONDS) diff --git a/main/manager-api/src/main/java/xiaozhi/modules/sys/controller/SysParamsController.java b/main/manager-api/src/main/java/xiaozhi/modules/sys/controller/SysParamsController.java index 24b070c0..26e6d245 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/sys/controller/SysParamsController.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/sys/controller/SysParamsController.java @@ -174,7 +174,7 @@ public class SysParamsController { return; } if (StringUtils.isBlank(url) || url.equals("null")) { - throw new RenException(ErrorCode.OTA_URL_EMPTY); + return; } // 检查是否包含localhost或127.0.0.1 @@ -211,7 +211,7 @@ public class SysParamsController { return; } if (StringUtils.isBlank(url) || url.equals("null")) { - throw new RenException(ErrorCode.MCP_URL_EMPTY); + return; } if (url.contains("localhost") || url.contains("127.0.0.1")) { throw new RenException(ErrorCode.MCP_URL_LOCALHOST); @@ -242,7 +242,7 @@ public class SysParamsController { return; } if (StringUtils.isBlank(url) || url.equals("null")) { - throw new RenException(ErrorCode.VOICEPRINT_URL_EMPTY); + return; } if (url.contains("localhost") || url.contains("127.0.0.1")) { throw new RenException(ErrorCode.VOICEPRINT_URL_LOCALHOST); diff --git a/main/manager-api/src/main/resources/db/changelog/202512031517.sql b/main/manager-api/src/main/resources/db/changelog/202512031517.sql new file mode 100644 index 00000000..2ad5bf22 --- /dev/null +++ b/main/manager-api/src/main/resources/db/changelog/202512031517.sql @@ -0,0 +1,6 @@ +-- 添加系统功能菜单配置参数 +delete from `sys_params` where param_code = 'system-web.menu'; + +-- 添加系统功能菜单配置参数 +INSERT INTO `sys_params` (id, param_code, param_value, value_type, param_type, remark) VALUES +(600, 'system-web.menu', '{"features":{"voiceprintRecognition":{"name":"feature.voiceprintRecognition.name","enabled":false,"description":"feature.voiceprintRecognition.description"},"voiceClone":{"name":"feature.voiceClone.name","enabled":false,"description":"feature.voiceClone.description"},"knowledgeBase":{"name":"feature.knowledgeBase.name","enabled":false,"description":"feature.knowledgeBase.description"},"mcpAccessPoint":{"name":"feature.mcpAccessPoint.name","enabled":false,"description":"feature.mcpAccessPoint.description"},"vad":{"name":"feature.vad.name","enabled":true,"description":"feature.vad.description"},"asr":{"name":"feature.asr.name","enabled":true,"description":"feature.asr.description"}},"groups":{"featureManagement":["voiceprintRecognition","voiceClone","knowledgeBase","mcpAccessPoint"],"voiceManagement":["vad","asr"]}}', 'json', 1, '系统功能菜单配置'); \ No newline at end of file diff --git a/main/manager-api/src/main/resources/db/changelog/202512041515.sql b/main/manager-api/src/main/resources/db/changelog/202512041515.sql new file mode 100644 index 00000000..5e26b99f --- /dev/null +++ b/main/manager-api/src/main/resources/db/changelog/202512041515.sql @@ -0,0 +1,14 @@ +-- liquibase formatted sql + +-- changeset xiaozhi:202512041515 +CREATE TABLE ai_agent_context_provider ( + id VARCHAR(32) NOT NULL COMMENT '主键', + agent_id VARCHAR(32) NOT NULL COMMENT '智能体ID', + context_providers JSON COMMENT '上下文源配置', + creator BIGINT COMMENT '创建者', + created_at DATETIME COMMENT '创建时间', + updater BIGINT COMMENT '更新者', + updated_at DATETIME COMMENT '更新时间', + PRIMARY KEY (id), + INDEX idx_agent_id (agent_id) +) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='智能体上下文源配置表'; diff --git a/main/manager-api/src/main/resources/db/changelog/202512131453.sql b/main/manager-api/src/main/resources/db/changelog/202512131453.sql new file mode 100644 index 00000000..27ea4a1b --- /dev/null +++ b/main/manager-api/src/main/resources/db/changelog/202512131453.sql @@ -0,0 +1,6 @@ +-- 删除server模块是否开启token认证参数 +delete from `sys_params` where param_code = 'server.auth.enabled'; + +-- 添加server模块是否开启token认证参数 +INSERT INTO `sys_params` (id, param_code, param_value, value_type, param_type, remark) VALUES +(122, 'server.auth.enabled', 'true', 'boolean', 1, 'server模块是否开启token认证'); \ No newline at end of file diff --git a/main/manager-api/src/main/resources/db/changelog/202512161529.sql b/main/manager-api/src/main/resources/db/changelog/202512161529.sql new file mode 100644 index 00000000..40f154c7 --- /dev/null +++ b/main/manager-api/src/main/resources/db/changelog/202512161529.sql @@ -0,0 +1 @@ +INSERT INTO `sys_params` (id, param_code, param_value, value_type, param_type, remark) VALUES (311, 'enable_websocket_ping', 'false', 'boolean', 1, '是否启用WebSocket心跳保活机制'); \ No newline at end of file diff --git a/main/manager-api/src/main/resources/db/changelog/202512192245.sql b/main/manager-api/src/main/resources/db/changelog/202512192245.sql new file mode 100644 index 00000000..7162fd1c --- /dev/null +++ b/main/manager-api/src/main/resources/db/changelog/202512192245.sql @@ -0,0 +1,2 @@ +-- 为智能体聊天历史记录添加音频ID索引 +ALTER TABLE ai_agent_chat_history ADD INDEX idx_ai_agent_chat_history_audio_id (audio_id); diff --git a/main/manager-api/src/main/resources/db/changelog/202512221117.sql b/main/manager-api/src/main/resources/db/changelog/202512221117.sql new file mode 100644 index 00000000..4a3c8d24 --- /dev/null +++ b/main/manager-api/src/main/resources/db/changelog/202512221117.sql @@ -0,0 +1,10 @@ +-- 更新豆包流式ASR供应器,增加end_window_size配置 +delete from `ai_model_provider` where id = 'SYSTEM_ASR_DoubaoStreamASR'; +INSERT INTO `ai_model_provider` (`id`, `model_type`, `provider_code`, `name`, `fields`, `sort`, `creator`, `create_date`, `updater`, `update_date`) VALUES +('SYSTEM_ASR_DoubaoStreamASR', 'ASR', 'doubao_stream', '火山引擎语音识别(流式)', '[{"key":"appid","label":"应用ID","type":"string"},{"key":"access_token","label":"访问令牌","type":"string"},{"key":"cluster","label":"集群","type":"string"},{"key":"boosting_table_name","label":"热词文件名称","type":"string"},{"key":"correct_table_name","label":"替换词文件名称","type":"string"},{"key":"output_dir","label":"输出目录","type":"string"},{"key":"end_window_size","label":"静音判定时长(ms)","type":"number"}]', 3, 1, NOW(), 1, NOW()); + + +-- 更新豆包流式ASR模型配置,增加end_window_size默认值 +UPDATE `ai_model_config` SET +`config_json` = JSON_SET(`config_json`, '$.end_window_size', 200) +WHERE `id` = 'ASR_DoubaoStreamASR' AND JSON_EXTRACT(`config_json`, '$.end_window_size') IS NULL; diff --git a/main/manager-api/src/main/resources/db/changelog/db.changelog-master.yaml b/main/manager-api/src/main/resources/db/changelog/db.changelog-master.yaml index 57c9ea07..00db89db 100755 --- a/main/manager-api/src/main/resources/db/changelog/db.changelog-master.yaml +++ b/main/manager-api/src/main/resources/db/changelog/db.changelog-master.yaml @@ -430,3 +430,46 @@ databaseChangeLog: - sqlFile: encoding: utf8 path: classpath:db/changelog/202511221450.sql + - changeSet: + id: 202512031517 + author: rainv123 + changes: + - sqlFile: + encoding: utf8 + path: classpath:db/changelog/202512031517.sql + + - changeSet: + id: 202512041515 + author: cgd + changes: + - sqlFile: + encoding: utf8 + path: classpath:db/changelog/202512041515.sql + - changeSet: + id: 202512131453 + author: hrz + changes: + - sqlFile: + encoding: utf8 + path: classpath:db/changelog/202512131453.sql + - changeSet: + id: 202512161529 + author: RanChen + changes: + - sqlFile: + encoding: utf8 + path: classpath:db/changelog/202512161529.sql + - changeSet: + id: 202512192245 + author: hrz + changes: + - sqlFile: + encoding: utf8 + path: classpath:db/changelog/202512192245.sql + - changeSet: + id: 202512221117 + author: RanChen + changes: + - sqlFile: + encoding: utf8 + path: classpath:db/changelog/202512221117.sql \ No newline at end of file diff --git a/main/manager-api/src/main/resources/mapper/agent/AiAgentChatHistoryDao.xml b/main/manager-api/src/main/resources/mapper/agent/AiAgentChatHistoryDao.xml index 3e80f7fd..84e22fd1 100644 --- a/main/manager-api/src/main/resources/mapper/agent/AiAgentChatHistoryDao.xml +++ b/main/manager-api/src/main/resources/mapper/agent/AiAgentChatHistoryDao.xml @@ -22,13 +22,18 @@ created_at, updated_at - - DELETE FROM ai_agent_chat_audio - WHERE id IN ( - SELECT audio_id - FROM ai_agent_chat_history - WHERE agent_id = #{agentId} - ) + + + + DELETE FROM ai_agent_chat_audio + WHERE id IN + + #{id} + diff --git a/main/manager-mobile/src/pages/settings/index.vue b/main/manager-mobile/src/pages/settings/index.vue index 90c4e6cd..b1fa487b 100644 --- a/main/manager-mobile/src/pages/settings/index.vue +++ b/main/manager-mobile/src/pages/settings/index.vue @@ -235,7 +235,7 @@ function showAbout() { title: t('settings.aboutApp', { appName: import.meta.env.VITE_APP_TITLE }), content: t('settings.aboutContent', { appName: import.meta.env.VITE_APP_TITLE, - version: '0.8.8' + version: '0.8.10' }), showCancel: false, confirmText: t('common.confirm'), diff --git a/main/manager-web/src/assets/xiaozhi-ai.png b/main/manager-web/src/assets/xiaozhi-ai.png index ef3834af..8a595493 100644 Binary files a/main/manager-web/src/assets/xiaozhi-ai.png and b/main/manager-web/src/assets/xiaozhi-ai.png differ diff --git a/main/manager-web/src/assets/xiaozhi-ai_de.png b/main/manager-web/src/assets/xiaozhi-ai_de.png new file mode 100644 index 00000000..f2c54d80 Binary files /dev/null and b/main/manager-web/src/assets/xiaozhi-ai_de.png differ diff --git a/main/manager-web/src/assets/xiaozhi-ai_en.png b/main/manager-web/src/assets/xiaozhi-ai_en.png new file mode 100644 index 00000000..d35fd499 Binary files /dev/null and b/main/manager-web/src/assets/xiaozhi-ai_en.png differ diff --git a/main/manager-web/src/assets/xiaozhi-ai_vi.png b/main/manager-web/src/assets/xiaozhi-ai_vi.png new file mode 100644 index 00000000..a54913ad Binary files /dev/null and b/main/manager-web/src/assets/xiaozhi-ai_vi.png differ diff --git a/main/manager-web/src/assets/xiaozhi-ai_zh_CN.png b/main/manager-web/src/assets/xiaozhi-ai_zh_CN.png new file mode 100644 index 00000000..ef3834af Binary files /dev/null and b/main/manager-web/src/assets/xiaozhi-ai_zh_CN.png differ diff --git a/main/manager-web/src/assets/xiaozhi-ai_zh_TW.png b/main/manager-web/src/assets/xiaozhi-ai_zh_TW.png new file mode 100644 index 00000000..ac5f59aa Binary files /dev/null and b/main/manager-web/src/assets/xiaozhi-ai_zh_TW.png differ diff --git a/main/manager-web/src/components/AddModelDialog.vue b/main/manager-web/src/components/AddModelDialog.vue index c95e95a9..862d8b29 100644 --- a/main/manager-web/src/components/AddModelDialog.vue +++ b/main/manager-web/src/components/AddModelDialog.vue @@ -404,7 +404,7 @@ export default { .custom-input-bg .el-input__inner, .custom-input-bg .el-textarea__inner { - background-color: #f6f8fc; + background-color: #ffffff; } diff --git a/main/manager-web/src/components/ContextProviderDialog.vue b/main/manager-web/src/components/ContextProviderDialog.vue new file mode 100644 index 00000000..84eae4a9 --- /dev/null +++ b/main/manager-web/src/components/ContextProviderDialog.vue @@ -0,0 +1,329 @@ + + + + + diff --git a/main/manager-web/src/components/DeviceItem.vue b/main/manager-web/src/components/DeviceItem.vue index bfbbb97f..c29d3243 100644 --- a/main/manager-web/src/components/DeviceItem.vue +++ b/main/manager-web/src/components/DeviceItem.vue @@ -23,7 +23,7 @@
{{ $t('home.configureRole') }}
-
+
{{ $t('home.voiceprintRecognition') }}
@@ -49,7 +49,15 @@ import i18n from '@/i18n'; export default { name: 'DeviceItem', props: { - device: { type: Object, required: true } + device: { type: Object, required: true }, + featureStatus: { + type: Object, + default: () => ({ + voiceprintRecognition: false, + voiceClone: false, + knowledgeBase: false + }) + } }, data() { return { switchValue: false } diff --git a/main/manager-web/src/components/FunctionDialog.vue b/main/manager-web/src/components/FunctionDialog.vue index 23adce78..23c4d1ad 100644 --- a/main/manager-web/src/components/FunctionDialog.vue +++ b/main/manager-web/src/components/FunctionDialog.vue @@ -106,7 +106,7 @@
-
+
@@ -171,6 +171,7 @@ + + \ No newline at end of file diff --git a/main/manager-web/src/views/KnowledgeFileUpload.vue b/main/manager-web/src/views/KnowledgeFileUpload.vue index 6d58ba2f..37fd7e27 100644 --- a/main/manager-web/src/views/KnowledgeFileUpload.vue +++ b/main/manager-web/src/views/KnowledgeFileUpload.vue @@ -214,7 +214,7 @@
- +
@@ -61,6 +62,7 @@ import ChatHistoryDialog from '@/components/ChatHistoryDialog.vue'; import DeviceItem from '@/components/DeviceItem.vue'; import HeaderBar from '@/components/HeaderBar.vue'; import VersionFooter from '@/components/VersionFooter.vue'; +import featureManager from '@/utils/featureManager'; export default { name: 'HomePage', @@ -76,15 +78,33 @@ export default { skeletonCount: localStorage.getItem('skeletonCount') || 8, showChatHistory: false, currentAgentId: '', - currentAgentName: '' + currentAgentName: '', + // 功能状态 + featureStatus: { + voiceprintRecognition: false, + voiceClone: false, + knowledgeBase: false + } } }, - mounted() { + async mounted() { this.fetchAgentList(); + await this.loadFeatureStatus(); }, methods: { + // 加载功能状态 + async loadFeatureStatus() { + await featureManager.waitForInitialization(); + const config = featureManager.getConfig(); + this.featureStatus = { + voiceprintRecognition: config.voiceprintRecognition, + voiceClone: config.voiceClone, + knowledgeBase: config.knowledgeBase + }; + }, + showAddDialog() { this.addDeviceDialogVisible = true }, diff --git a/main/manager-web/src/views/login.vue b/main/manager-web/src/views/login.vue index 9d56fdc5..f9dd0748 100644 --- a/main/manager-web/src/views/login.vue +++ b/main/manager-web/src/views/login.vue @@ -10,7 +10,7 @@ gap: 10px; "> - +
+ +
+ + {{ $t('roleConfig.contextProviderSuccess', { count: currentContextProviders.length }) }}{{ $t('roleConfig.contextProviderDocLink') }} + + + {{ $t('roleConfig.editContextProvider') }} + +
+
- +
- +
+
@@ -275,14 +302,17 @@ import Api from "@/apis/api"; import { getServiceUrl } from "@/apis/api"; import RequestService from "@/apis/httpRequest"; import FunctionDialog from "@/components/FunctionDialog.vue"; +import ContextProviderDialog from "@/components/ContextProviderDialog.vue"; import HeaderBar from "@/components/HeaderBar.vue"; import i18n from "@/i18n"; +import featureManager from "@/utils/featureManager"; export default { name: "RoleConfigPage", - components: { HeaderBar, FunctionDialog }, + components: { HeaderBar, FunctionDialog, ContextProviderDialog }, data() { return { + showContextProviderDialog: false, form: { agentCode: "", agentName: "", @@ -320,12 +350,18 @@ export default { voiceDetails: {}, // 保存完整的音色信息 showFunctionDialog: false, currentFunctions: [], + currentContextProviders: [], allFunctions: [], originalFunctions: [], playingVoice: false, isPaused: false, currentAudio: null, currentPlayingVoiceId: null, + // 功能状态 + featureStatus: { + vad: false, // 语言检测活动功能状态 + asr: false, // 语音识别功能状态 + }, }; }, methods: { @@ -356,6 +392,7 @@ export default { paramInfo: item.params, }; }), + contextProviders: this.currentContextProviders, }; Api.agent.updateAgentConfig(this.$route.query.agentId, configData, ({ data }) => { if (data.code === 0) { @@ -472,6 +509,9 @@ export default { }; // 后端只给了最小映射:[{ id, agentId, pluginId }, ...] const savedMappings = data.data.functions || []; + + // 加载上下文配置 + this.currentContextProviders = data.data.contextProviders || []; // 先保证 allFunctions 已经加载(如果没有,则先 fetchAllFunctions) const ensureFuncs = this.allFunctions.length @@ -646,6 +686,12 @@ export default { this.showFunctionDialog = true; } }, + openContextProviderDialog() { + this.showContextProviderDialog = true; + }, + handleUpdateContext(providers) { + this.currentContextProviders = providers; + }, handleUpdateFunctions(selected) { this.currentFunctions = selected; }, @@ -980,6 +1026,19 @@ export default { this.form.chatHistoryConf = 0; } }, + // 加载功能状态 + async loadFeatureStatus() { + try { + // 确保featureManager已初始化完成 + await featureManager.waitForInitialization(); + const config = featureManager.getConfig(); + this.featureStatus.voiceprintRecognition = config.voiceprintRecognition || false; + this.featureStatus.vad = config.vad || false; + this.featureStatus.asr = config.asr || false; + } catch (error) { + console.error("加载功能状态失败:", error); + } + }, }, watch: { "form.model.ttsModelId": { @@ -1002,7 +1061,7 @@ export default { immediate: true, }, }, - mounted() { + async mounted() { const agentId = this.$route.query.agentId; if (agentId) { this.fetchAgentConfig(agentId); @@ -1010,6 +1069,8 @@ export default { } this.fetchModelOptions(); this.fetchTemplates(); + // 加载功能状态,确保featureManager已初始化 + await this.loadFeatureStatus(); }, }; @@ -1298,6 +1359,26 @@ export default { justify-content: flex-end; } +.chat-history-options ::v-deep .el-radio-button { + border-color: #5778ff; +} + +.chat-history-options ::v-deep .el-radio-button .el-radio-button__inner { + color: #5778ff; + border-color: #5778ff; + background-color: transparent; +} + +.chat-history-options ::v-deep .el-radio-button.is-active .el-radio-button__inner { + background-color: #5778ff; + border-color: #5778ff; + color: white; +} + +.chat-history-options ::v-deep .el-radio-button .el-radio-button__inner:hover { + color: #5778ff; +} + .header-actions { display: flex; align-items: center; @@ -1345,4 +1426,18 @@ export default { height: 32px; margin-left: 8px; } + +.context-provider-item ::v-deep .el-form-item__label { + line-height: 42px !important; +} + +.doc-link { + color: #5778ff; + text-decoration: none; + margin-left: 4px; + + &:hover { + text-decoration: underline; + } +} diff --git a/main/xiaozhi-server/agent-base-prompt.txt b/main/xiaozhi-server/agent-base-prompt.txt index 490b23d3..7fd90d21 100644 --- a/main/xiaozhi-server/agent-base-prompt.txt +++ b/main/xiaozhi-server/agent-base-prompt.txt @@ -74,6 +74,7 @@ - **今天农历:** {{lunar_date}} - **用户所在城市:** {{local_address}} - **当地未来7天天气:** {{weather_info}} +{{ dynamic_context }} diff --git a/main/xiaozhi-server/app.py b/main/xiaozhi-server/app.py index bc6f0b54..ab085464 100644 --- a/main/xiaozhi-server/app.py +++ b/main/xiaozhi-server/app.py @@ -9,6 +9,7 @@ from core.utils.util import get_local_ip, validate_mcp_endpoint from core.http_server import SimpleHttpServer from core.websocket_server import WebSocketServer from core.utils.util import check_ffmpeg_installed +from core.utils.gc_manager import get_gc_manager TAG = __name__ logger = setup_logging() @@ -63,6 +64,10 @@ async def main(): # 添加 stdin 监控任务 stdin_task = asyncio.create_task(monitor_stdin()) + # 启动全局GC管理器(5分钟清理一次) + gc_manager = get_gc_manager(interval_seconds=300) + await gc_manager.start() + # 启动 WebSocket 服务器 ws_server = WebSocketServer(config) ws_task = asyncio.create_task(ws_server.start()) @@ -122,6 +127,9 @@ async def main(): except asyncio.CancelledError: print("任务被取消,清理资源中...") finally: + # 停止全局GC管理器 + await gc_manager.stop() + # 取消所有任务(关键修复点) stdin_task.cancel() ws_task.cancel() diff --git a/main/xiaozhi-server/config.yaml b/main/xiaozhi-server/config.yaml index de804ac1..1139dbf4 100644 --- a/main/xiaozhi-server/config.yaml +++ b/main/xiaozhi-server/config.yaml @@ -69,6 +69,9 @@ enable_greeting: true enable_stop_tts_notify: false # 说完话是否开启提示音,音效地址 stop_tts_notify_voice: "config/assets/tts_notify.mp3" +# 是否启用WebSocket心跳保活机制 +enable_websocket_ping: false + # TTS音频发送延迟配置 # tts_audio_send_delay: 控制音频包发送间隔 @@ -113,6 +116,15 @@ wakeup_words: # MCP接入点地址,地址格式为:ws://你的mcp接入点ip或者域名:端口号/mcp/?token=你的token # 详细教程 https://github.com/xinnan-tech/xiaozhi-esp32-server/blob/main/docs/mcp-endpoint-integration.md mcp_endpoint: 你的接入点 websocket地址 + +# 上下文源配置 +# 用于在系统提示词中注入动态数据,如健康数据、股票信息等 +# 可以添加多个上下文源 +context_providers: + - url: "" + headers: + Authorization: "" + # 插件的基础配置 plugins: # 获取天气插件的配置,这里填写你的api_key @@ -334,6 +346,8 @@ ASR: # 热词、替换词使用流程:https://www.volcengine.com/docs/6561/155738 boosting_table_name: (选填)你的热词文件名称 correct_table_name: (选填)你的替换词文件名称 + # 静音判定时长(ms),默认200ms + end_window_size: 200 output_dir: tmp/ TencentASR: # token申请地址:https://console.cloud.tencent.com/cam/capi @@ -462,7 +476,6 @@ ASR: domain: slm # 识别领域,iat:日常用语,medical:医疗,finance:金融等 language: zh_cn # 语言,zh_cn:中文,en_us:英文 accent: mandarin # 方言,mandarin:普通话 - dwa: wpgs # 动态修正,wpgs:实时返回中间结果 # 调整音频处理参数以提高长语音识别质量 output_dir: tmp/ diff --git a/main/xiaozhi-server/config/config_loader.py b/main/xiaozhi-server/config/config_loader.py index b70cf961..e0220e33 100644 --- a/main/xiaozhi-server/config/config_loader.py +++ b/main/xiaozhi-server/config/config_loader.py @@ -32,7 +32,16 @@ def load_config(): custom_config = read_config(custom_config_path) if custom_config.get("manager-api", {}).get("url"): - config = get_config_from_api(custom_config) + import asyncio + try: + loop = asyncio.get_running_loop() + # 如果已经在事件循环中,使用异步版本 + config = asyncio.run_coroutine_threadsafe( + get_config_from_api_async(custom_config), loop + ).result() + except RuntimeError: + # 如果不在事件循环中(启动时),创建新的事件循环 + config = asyncio.run(get_config_from_api_async(custom_config)) else: # 合并配置 config = merge_configs(default_config, custom_config) @@ -44,13 +53,13 @@ def load_config(): return config -def get_config_from_api(config): - """从Java API获取配置""" +async def get_config_from_api_async(config): + """从Java API获取配置(异步版本)""" # 初始化API客户端 init_service(config) # 获取服务器配置 - config_data = get_server_config() + config_data = await get_server_config() if config_data is None: raise Exception("Failed to fetch server config from API") @@ -59,6 +68,7 @@ def get_config_from_api(config): "url": config["manager-api"].get("url", ""), "secret": config["manager-api"].get("secret", ""), } + auth_enabled = config_data.get("server", {}).get("auth", {}).get("enabled", False) # server的配置以本地为准 if config.get("server"): config_data["server"] = { @@ -68,15 +78,16 @@ def get_config_from_api(config): "vision_explain": config["server"].get("vision_explain", ""), "auth_key": config["server"].get("auth_key", ""), } + config_data["server"]["auth"] = {"enabled": auth_enabled} # 如果服务器没有prompt_template,则从本地配置读取 if not config_data.get("prompt_template"): config_data["prompt_template"] = config.get("prompt_template") return config_data -def get_private_config_from_api(config, device_id, client_id): +async def get_private_config_from_api(config, device_id, client_id): """从Java API获取私有配置""" - return get_agent_models(device_id, client_id, config["selected_module"]) + return await get_agent_models(device_id, client_id, config["selected_module"]) def ensure_directories(config): diff --git a/main/xiaozhi-server/config/logger.py b/main/xiaozhi-server/config/logger.py index 37e5eabf..bf64893c 100644 --- a/main/xiaozhi-server/config/logger.py +++ b/main/xiaozhi-server/config/logger.py @@ -5,7 +5,7 @@ from config.config_loader import load_config from config.settings import check_config_file from datetime import datetime -SERVER_VERSION = "0.8.8" +SERVER_VERSION = "0.8.10" _logger_initialized = False diff --git a/main/xiaozhi-server/config/manage_api_client.py b/main/xiaozhi-server/config/manage_api_client.py index cd3f86c4..f899d61c 100644 --- a/main/xiaozhi-server/config/manage_api_client.py +++ b/main/xiaozhi-server/config/manage_api_client.py @@ -1,5 +1,4 @@ import os -import time import base64 from typing import Optional, Dict @@ -20,7 +19,7 @@ class DeviceBindException(Exception): class ManageApiClient: _instance = None - _client = None + _async_clients = {} # 为每个事件循环存储独立的客户端 _secret = None def __new__(cls, config): @@ -32,7 +31,7 @@ class ManageApiClient: @classmethod def _init_client(cls, config): - """初始化持久化连接池""" + """初始化配置(延迟创建客户端)""" cls.config = config.get("manager-api") if not cls.config: @@ -47,23 +46,41 @@ class ManageApiClient: cls._secret = cls.config.get("secret") cls.max_retries = cls.config.get("max_retries", 6) # 最大重试次数 cls.retry_delay = cls.config.get("retry_delay", 10) # 初始重试延迟(秒) - # NOTE(goody): 2025/4/16 http相关资源统一管理,后续可以增加线程池或者超时 - # 后续也可以统一配置apiToken之类的走通用的Auth - cls._client = httpx.Client( - base_url=cls.config.get("url"), - headers={ - "User-Agent": f"PythonClient/2.0 (PID:{os.getpid()})", - "Accept": "application/json", - "Authorization": "Bearer " + cls._secret, - }, - timeout=cls.config.get("timeout", 30), # 默认超时时间30秒 - ) + # 不在这里创建 AsyncClient,延迟到实际使用时创建 + cls._async_clients = {} @classmethod - def _request(cls, method: str, endpoint: str, **kwargs) -> Dict: - """发送单次HTTP请求并处理响应""" + async def _ensure_async_client(cls): + """确保异步客户端已创建(为每个事件循环创建独立的客户端)""" + import asyncio + + try: + loop = asyncio.get_running_loop() + loop_id = id(loop) + + # 为每个事件循环创建独立的客户端 + if loop_id not in cls._async_clients: + cls._async_clients[loop_id] = httpx.AsyncClient( + base_url=cls.config.get("url"), + headers={ + "User-Agent": f"PythonClient/2.0 (PID:{os.getpid()})", + "Accept": "application/json", + "Authorization": "Bearer " + cls._secret, + }, + timeout=cls.config.get("timeout", 30), + ) + return cls._async_clients[loop_id] + except RuntimeError: + # 如果没有运行中的事件循环,创建一个临时的 + raise Exception("必须在异步上下文中调用") + + @classmethod + async def _async_request(cls, method: str, endpoint: str, **kwargs) -> Dict: + """发送单次异步HTTP请求并处理响应""" + # 确保客户端已创建 + client = await cls._ensure_async_client() endpoint = endpoint.lstrip("/") - response = cls._client.request(method, endpoint, **kwargs) + response = await client.request(method, endpoint, **kwargs) response.raise_for_status() result = response.json() @@ -96,22 +113,24 @@ class ManageApiClient: return False @classmethod - def _execute_request(cls, method: str, endpoint: str, **kwargs) -> Dict: - """带重试机制的请求执行器""" + async def _execute_async_request(cls, method: str, endpoint: str, **kwargs) -> Dict: + """带重试机制的异步请求执行器""" + import asyncio + retry_count = 0 while retry_count <= cls.max_retries: try: - # 执行请求 - return cls._request(method, endpoint, **kwargs) + # 执行异步请求 + return await cls._async_request(method, endpoint, **kwargs) except Exception as e: # 判断是否应该重试 if retry_count < cls.max_retries and cls._should_retry(e): retry_count += 1 print( - f"{method} {endpoint} 请求失败,将在 {cls.retry_delay:.1f} 秒后进行第 {retry_count} 次重试" + f"{method} {endpoint} 异步请求失败,将在 {cls.retry_delay:.1f} 秒后进行第 {retry_count} 次重试" ) - time.sleep(cls.retry_delay) + await asyncio.sleep(cls.retry_delay) continue else: # 不重试,直接抛出异常 @@ -119,22 +138,30 @@ class ManageApiClient: @classmethod def safe_close(cls): - """安全关闭连接池""" - if cls._client: - cls._client.close() - cls._instance = None + """安全关闭所有异步连接池""" + import asyncio + + for client in list(cls._async_clients.values()): + try: + asyncio.run(client.aclose()) + except Exception: + pass + cls._async_clients.clear() + cls._instance = None -def get_server_config() -> Optional[Dict]: +async def get_server_config() -> Optional[Dict]: """获取服务器基础配置""" - return ManageApiClient._instance._execute_request("POST", "/config/server-base") + return await ManageApiClient._instance._execute_async_request( + "POST", "/config/server-base" + ) -def get_agent_models( +async def get_agent_models( mac_address: str, client_id: str, selected_module: Dict ) -> Optional[Dict]: """获取代理模型配置""" - return ManageApiClient._instance._execute_request( + return await ManageApiClient._instance._execute_async_request( "POST", "/config/agent-models", json={ @@ -145,28 +172,26 @@ def get_agent_models( ) -def save_mem_local_short(mac_address: str, short_momery: str) -> Optional[Dict]: +async def generate_and_save_chat_summary(session_id: str) -> Optional[Dict]: + """生成并保存聊天记录总结""" try: - return ManageApiClient._instance._execute_request( - "PUT", - f"/agent/saveMemory/" + mac_address, - json={ - "summaryMemory": short_momery, - }, + return await ManageApiClient._instance._execute_async_request( + "POST", + f"/agent/chat-summary/{session_id}/save", ) except Exception as e: - print(f"存储短期记忆到服务器失败: {e}") + print(f"生成并保存聊天记录总结失败: {e}") return None -def report( +async def report( mac_address: str, session_id: str, chat_type: int, content: str, audio, report_time ) -> Optional[Dict]: - """带熔断的业务方法示例""" + """异步聊天记录上报""" if not content or not ManageApiClient._instance: return None try: - return ManageApiClient._instance._execute_request( + return await ManageApiClient._instance._execute_async_request( "POST", f"/agent/chat-history/report", json={ diff --git a/main/xiaozhi-server/core/api/base_handler.py b/main/xiaozhi-server/core/api/base_handler.py index 7330185e..db277543 100644 --- a/main/xiaozhi-server/core/api/base_handler.py +++ b/main/xiaozhi-server/core/api/base_handler.py @@ -10,7 +10,15 @@ class BaseHandler: def _add_cors_headers(self, response): """添加CORS头信息""" response.headers["Access-Control-Allow-Headers"] = ( - "client-id, content-type, device-id" + "client-id, content-type, device-id, authorization" ) response.headers["Access-Control-Allow-Credentials"] = "true" response.headers["Access-Control-Allow-Origin"] = "*" + + async def handle_options(self, request): + """处理OPTIONS请求,添加CORS头信息""" + response = web.Response(body=b"", content_type="text/plain") + self._add_cors_headers(response) + # 添加允许的方法 + response.headers["Access-Control-Allow-Methods"] = "GET, POST, OPTIONS" + return response diff --git a/main/xiaozhi-server/core/api/ota_handler.py b/main/xiaozhi-server/core/api/ota_handler.py index b6c88dff..1e3b3bd0 100644 --- a/main/xiaozhi-server/core/api/ota_handler.py +++ b/main/xiaozhi-server/core/api/ota_handler.py @@ -3,15 +3,46 @@ import time import base64 import hashlib import hmac +import os +import re +import glob +from typing import Dict, List, Tuple from aiohttp import web from core.auth import AuthManager -from core.utils.util import get_local_ip +from core.utils.util import get_local_ip, get_vision_url from core.api.base_handler import BaseHandler TAG = __name__ +def _safe_basename(filename: str) -> str: + # Prevent directory traversal + return os.path.basename(filename) + + +def _parse_version(ver: str) -> Tuple[int, ...]: + # conservative parser: split by non-digit, keep numeric parts + parts = re.findall(r"\d+", ver) + return tuple(int(p) for p in parts) if parts else (0,) + + +def _is_higher_version(a: str, b: str) -> bool: + """Return True if version string a > b (semver-like numeric compare).""" + ta = _parse_version(a) + tb = _parse_version(b) + # compare tuple lexicographically, but allow different lengths + maxlen = max(len(ta), len(tb)) + for i in range(maxlen): + ai = ta[i] if i < len(ta) else 0 + bi = tb[i] if i < len(tb) else 0 + if ai > bi: + return True + if ai < bi: + return False + return False + + class OTAHandler(BaseHandler): def __init__(self, config: dict): super().__init__(config) @@ -23,6 +54,54 @@ class OTAHandler(BaseHandler): expire_seconds = auth_config.get("expire_seconds") self.auth = AuthManager(secret_key=secret_key, expire_seconds=expire_seconds) + # firmware storage + self.bin_dir = os.path.join(os.getcwd(), "data", "bin") + # cache structure: { 'updated_at': timestamp, 'ttl': seconds, 'files_by_model': { model: [(version, filename), ...] } } + self._bin_cache: Dict = { + "updated_at": 0, + "ttl": config.get("firmware_cache_ttl", 30), + "files_by_model": {}, + } + + def _refresh_bin_cache_if_needed(self): + now = int(time.time()) + ttl = int(self._bin_cache.get("ttl", 30)) + if now - int( + self._bin_cache.get("updated_at", 0) + ) < ttl and self._bin_cache.get("files_by_model"): + return + + files_by_model: Dict[str, List[Tuple[str, str]]] = {} + try: + if not os.path.isdir(self.bin_dir): + os.makedirs(self.bin_dir, exist_ok=True) + + # match files like model_1.2.3.bin (allow dots, dashes, underscores in model and version) + pattern = os.path.join(self.bin_dir, "*.bin") + for path in glob.glob(pattern): + fname = os.path.basename(path) + # filename format: {model}_{version}.bin + m = re.match(r"^(.+?)_([0-9][A-Za-z0-9\.\-_]*)\.bin$", fname) + if not m: + # skip files not conforming to naming rule + continue + model = m.group(1) + version = m.group(2) + files_by_model.setdefault(model, []).append((version, fname)) + + # sort versions for each model descending + for model, items in files_by_model.items(): + items.sort(key=lambda it: _parse_version(it[0]), reverse=True) + + self._bin_cache["files_by_model"] = files_by_model + self._bin_cache["updated_at"] = now + self.logger.bind(tag=TAG).info( + f"Firmware cache refreshed: {len(files_by_model)} models" + ) + except Exception as e: + self.logger.bind(tag=TAG).error(f"刷新固件缓存失败: {e}") + # keep previous cache if any + def generate_password_signature(self, content: str, secret_key: str) -> str: """生成MQTT密码签名 @@ -62,7 +141,14 @@ class OTAHandler(BaseHandler): return f"ws://{local_ip}:{port}/xiaozhi/v1/" async def handle_post(self, request): - """处理 OTA POST 请求""" + """处理 OTA POST 请求 + + This handler will: + - read device id/client id (as before) + - attempt to determine device model and current firmware version (prefer headers, fallback to body) + - check data/bin for newer firmware for that model + - if found a newer firmware, set firmware.url to the download endpoint + """ try: data = await request.text() self.logger.bind(tag=TAG).debug(f"OTA请求方法: {request.method}") @@ -81,33 +167,76 @@ class OTAHandler(BaseHandler): else: raise Exception("OTA请求ClientID为空") - data_json = json.loads(data) + data_json = {} + try: + data_json = json.loads(data) if data else {} + except Exception: + data_json = {} server_config = self.config["server"] - port = int(server_config.get("port", 8000)) + # Distinguish ports: + # - websocket_port is used to construct websocket URL (server["port"]) + # - http_port is used to construct OTA download URLs (server["http_port"]) + websocket_port = int(server_config.get("port", 8000)) + http_port = int(server_config.get("http_port", 8003)) local_ip = get_local_ip() + # Determine device model (prefer headers) + device_model = "" + # header candidates + for h in ("device-model", "device_model", "model"): + if h in request.headers: + device_model = request.headers.get(h, "").strip() + break + # body fallback + if not device_model: + try: + if "board" in data_json and isinstance(data_json["board"], dict): + device_model = data_json["board"].get("type", "") + elif "model" in data_json: + device_model = data_json.get("model", "") + except Exception: + device_model = "" + if not device_model: + device_model = "default" + + # Determine device current version (prefer headers) + device_version = "" + for h in ( + "device-version", + "device_version", + "firmware-version", + "app-version", + "application-version", + ): + if h in request.headers: + device_version = request.headers.get(h, "").strip() + break + if not device_version: + try: + device_version = data_json.get("application", {}).get("version", "") + except Exception: + device_version = "" + if not device_version: + device_version = "0.0.0" + return_json = { "server_time": { "timestamp": int(round(time.time() * 1000)), "timezone_offset": server_config.get("timezone_offset", 8) * 60, }, "firmware": { - "version": data_json["application"].get("version", "1.0.0"), + "version": device_version, "url": "", }, } + # existing mqtt/websocket logic (unchanged) mqtt_gateway_endpoint = server_config.get("mqtt_gateway") if mqtt_gateway_endpoint: # 如果配置了非空字符串 - # 尝试从请求数据中获取设备型号 - device_model = "default" + # 尝试从请求数据中获取设备型号(已解析 above) try: - if "device" in data_json and isinstance(data_json["device"], dict): - device_model = data_json["device"].get("model", "default") - elif "model" in data_json: - device_model = data_json["model"] group_id = f"GID_{device_model}".replace(":", "_").replace(" ", "_") except Exception as e: self.logger.bind(tag=TAG).error(f"获取设备型号失败: {e}") @@ -159,20 +288,61 @@ class OTAHandler(BaseHandler): token = self.auth.generate_token(client_id, device_id) else: token = self.auth.generate_token(client_id, device_id) + # NOTE: use websocket_port here return_json["websocket"] = { - "url": self._get_websocket_url(local_ip, port), + "url": self._get_websocket_url(local_ip, websocket_port), "token": token, } self.logger.bind(tag=TAG).info( f"未配置MQTT网关,为设备 {device_id} 下发WebSocket配置" ) - self.logger.bind(tag=TAG).info(f"{return_json}") + + # Now check firmware files for updates + try: + self._refresh_bin_cache_if_needed() + files_by_model = self._bin_cache.get("files_by_model", {}) + candidates = files_by_model.get(device_model, []) + + self.logger.bind(tag=TAG).info( + f"查找型号 {device_model} 的固件,找到 {len(candidates)} 个候选" + ) + + chosen_url = "" + chosen_version = device_version + + # candidates are sorted descending by version + for ver, fname in candidates: + if _is_higher_version(ver, device_version): + # build download url (only allow download via our download endpoint) + chosen_version = ver + # Use get_vision_url to get the base URL and replace the path + vision_url = get_vision_url(self.config) + # Replace the path from "/mcp/vision/explain" to "/xiaozhi/ota/download/{fname}" + chosen_url = vision_url.replace( + "/mcp/vision/explain", f"/xiaozhi/ota/download/{fname}" + ) + break + + if chosen_url: + return_json["firmware"]["version"] = chosen_version + return_json["firmware"]["url"] = chosen_url + self.logger.bind(tag=TAG).info( + f"为设备 {device_id} 下发固件 {chosen_version} [如果地址前缀有误,请检查配置文件中的server.vision_explain]-> {chosen_url} " + ) + else: + self.logger.bind(tag=TAG).info( + f"设备 {device_id} 固件已是最新: {device_version}" + ) + + except Exception as e: + self.logger.bind(tag=TAG).error(f"检查固件版本时出错: {e}") response = web.Response( text=json.dumps(return_json, separators=(",", ":")), content_type="application/json", ) except Exception as e: + self.logger.bind(tag=TAG).error(f"OTA POST处理异常: {e}") return_json = {"success": False, "message": "request error."} response = web.Response( text=json.dumps(return_json, separators=(",", ":")), @@ -187,8 +357,9 @@ class OTAHandler(BaseHandler): try: server_config = self.config["server"] local_ip = get_local_ip() - port = int(server_config.get("port", 8000)) - websocket_url = self._get_websocket_url(local_ip, port) + # use websocket port for websocket URL + websocket_port = int(server_config.get("port", 8000)) + websocket_url = self._get_websocket_url(local_ip, websocket_port) message = f"OTA接口运行正常,向设备发送的websocket地址是:{websocket_url}" response = web.Response(text=message, content_type="text/plain") except Exception as e: @@ -197,3 +368,48 @@ class OTAHandler(BaseHandler): finally: self._add_cors_headers(response) return response + + async def handle_download(self, request): + """ + 下载固件接口 + URL: /xiaozhi/ota/download/{filename} + - 只允许下载 data/bin 目录下的 .bin 文件 + - filename 必须是 basename 且匹配安全的模式 + """ + try: + fname = request.match_info.get("filename", "") + if not fname: + raise web.HTTPBadRequest(text="filename required") + + # sanitize + fname = _safe_basename(fname) + # pattern: allow letters, numbers, dot, underscore, dash + if not re.match(r"^[A-Za-z0-9\.\-_]+\.bin$", fname): + raise web.HTTPBadRequest(text="invalid filename") + + file_path = os.path.join(self.bin_dir, fname) + # ensure realpath is under bin_dir + file_real = os.path.realpath(file_path) + bin_dir_real = os.path.realpath(self.bin_dir) + if ( + not file_real.startswith(bin_dir_real + os.sep) + and file_real != bin_dir_real + ): + raise web.HTTPForbidden(text="forbidden") + + if not os.path.isfile(file_real): + raise web.HTTPNotFound(text="file not found") + + # use FileResponse to stream file + resp = web.FileResponse(path=file_real) + except web.HTTPError as e: + resp = e + except Exception as e: + self.logger.bind(tag=TAG).error(f"固件下载异常: {e}") + resp = web.Response(text="download error", status=500) + finally: + try: + self._add_cors_headers(resp) + except Exception: + pass + return resp diff --git a/main/xiaozhi-server/core/api/vision_handler.py b/main/xiaozhi-server/core/api/vision_handler.py index 70be247a..28f48753 100644 --- a/main/xiaozhi-server/core/api/vision_handler.py +++ b/main/xiaozhi-server/core/api/vision_handler.py @@ -2,6 +2,7 @@ import json import copy from aiohttp import web from config.logger import setup_logging +from core.api.base_handler import BaseHandler from core.utils.util import get_vision_url, is_valid_image_file from core.utils.vllm import create_instance from config.config_loader import get_private_config_from_api @@ -16,10 +17,9 @@ TAG = __name__ MAX_FILE_SIZE = 5 * 1024 * 1024 -class VisionHandler: +class VisionHandler(BaseHandler): def __init__(self, config: dict): - self.config = config - self.logger = setup_logging() + super().__init__(config) # 初始化认证工具 self.auth = AuthToken(config["server"]["auth_key"]) @@ -96,7 +96,7 @@ class VisionHandler: current_config = copy.deepcopy(self.config) read_config_from_api = current_config.get("read_config_from_api", False) if read_config_from_api: - current_config = get_private_config_from_api( + current_config = await get_private_config_from_api( current_config, device_id, client_id, @@ -172,11 +172,3 @@ class VisionHandler: finally: self._add_cors_headers(response) return response - - def _add_cors_headers(self, response): - """添加CORS头信息""" - response.headers["Access-Control-Allow-Headers"] = ( - "client-id, content-type, device-id" - ) - response.headers["Access-Control-Allow-Credentials"] = "true" - response.headers["Access-Control-Allow-Origin"] = "*" diff --git a/main/xiaozhi-server/core/connection.py b/main/xiaozhi-server/core/connection.py index d2891393..ad8b4385 100644 --- a/main/xiaozhi-server/core/connection.py +++ b/main/xiaozhi-server/core/connection.py @@ -68,8 +68,12 @@ class ConnectionHandler: self.logger = setup_logging() self.server = server # 保存server实例的引用 - self.need_bind = False - self.bind_code = None + self.need_bind = False # 是否需要绑定设备 + self.bind_completed_event = asyncio.Event() + self.bind_code = None # 绑定设备的验证码 + self.last_bind_prompt_time = 0 # 上次播放绑定提示的时间戳(秒) + self.bind_prompt_interval = 60 # 绑定提示播放间隔(秒) + self.read_config_from_api = self.config.get("read_config_from_api", False) self.websocket = None @@ -88,7 +92,7 @@ class ConnectionHandler: self.client_listen_mode = "auto" # 线程任务相关 - self.loop = asyncio.get_event_loop() + self.loop = None # 在 handle_connection 中获取运行中的事件循环 self.stop_event = threading.Event() self.executor = ThreadPoolExecutor(max_workers=5) @@ -116,6 +120,7 @@ class ConnectionHandler: self.client_audio_buffer = bytearray() self.client_have_voice = False self.client_voice_window = deque(maxlen=5) + self.first_activity_time = 0.0 # 记录首次活动的时间(毫秒) self.last_activity_time = 0.0 # 统一的活动时间戳(毫秒) self.client_voice_stop = False self.last_is_voice = False @@ -158,10 +163,13 @@ class ConnectionHandler: self.conn_from_mqtt_gateway = False # 初始化提示词管理器 - self.prompt_manager = PromptManager(config, self.logger) + self.prompt_manager = PromptManager(self.config, self.logger) async def handle_connection(self, ws): try: + # 获取运行中的事件循环(必须在异步上下文中) + self.loop = asyncio.get_running_loop() + # 获取并验证headers self.headers = dict(ws.request.headers) real_ip = self.headers.get("x-real-ip") or self.headers.get( @@ -187,6 +195,7 @@ class ConnectionHandler: self.logger.bind(tag=TAG).info("连接来自:MQTT网关") # 初始化活动时间戳 + self.first_activity_time = time.time() * 1000 self.last_activity_time = time.time() * 1000 # 启动超时检查任务 @@ -195,10 +204,8 @@ class ConnectionHandler: self.welcome_msg = self.config["xiaozhi"] self.welcome_msg["session_id"] = self.session_id - # 获取差异化配置 - self._initialize_private_config() - # 异步初始化 - self.executor.submit(self._initialize_components) + # 在后台初始化配置和组件(完全不阻塞主循环) + asyncio.create_task(self._background_initialize()) try: async for message in self.websocket: @@ -237,7 +244,9 @@ class ConnectionHandler: loop = asyncio.new_event_loop() asyncio.set_event_loop(loop) loop.run_until_complete( - self.memory.save_memory(self.dialogue.dialogue) + self.memory.save_memory( + self.dialogue.dialogue, self.session_id + ) ) except Exception as e: self.logger.bind(tag=TAG).error(f"保存记忆失败: {e}") @@ -260,8 +269,37 @@ class ConnectionHandler: f"保存记忆后关闭连接失败: {close_error}" ) + async def _discard_message_with_bind_prompt(self): + """丢弃消息并检查是否需要播放绑定提示""" + current_time = time.time() + # 检查是否需要播放绑定提示 + if current_time - self.last_bind_prompt_time >= self.bind_prompt_interval: + self.last_bind_prompt_time = current_time + # 复用现有的绑定提示逻辑 + from core.handle.receiveAudioHandle import check_bind_device + + asyncio.create_task(check_bind_device(self)) + async def _route_message(self, message): """消息路由""" + # 检查是否已经获取到真实的绑定状态 + if not self.bind_completed_event.is_set(): + # 还没有获取到真实状态,等待直到获取到真实状态或超时 + try: + await asyncio.wait_for(self.bind_completed_event.wait(), timeout=1) + except asyncio.TimeoutError: + # 超时仍未获取到真实状态,丢弃消息 + await self._discard_message_with_bind_prompt() + return + + # 已经获取到真实状态,检查是否需要绑定 + if self.need_bind: + # 需要绑定,丢弃消息 + await self._discard_message_with_bind_prompt() + return + + # 不需要绑定,继续处理消息 + if isinstance(message, str): await handleTextMessage(self, message) elif isinstance(message, bytes): @@ -391,6 +429,15 @@ class ConnectionHandler: def _initialize_components(self): try: + if self.tts is None: + self.tts = self._initialize_tts() + # 打开语音合成通道 + asyncio.run_coroutine_threadsafe( + self.tts.open_audio_channels(self), self.loop + ) + if self.need_bind: + self.bind_completed_event.set() + return self.selected_module_str = build_module_string( self.config.get("selected_module", {}) ) @@ -414,17 +461,10 @@ class ConnectionHandler: # 初始化声纹识别 self._initialize_voiceprint() - # 打开语音识别通道 asyncio.run_coroutine_threadsafe( self.asr.open_audio_channels(self), self.loop ) - if self.tts is None: - self.tts = self._initialize_tts() - # 打开语音合成通道 - asyncio.run_coroutine_threadsafe( - self.tts.open_audio_channels(self), self.loop - ) """加载记忆""" self._initialize_memory() @@ -439,6 +479,7 @@ class ConnectionHandler: self.logger.bind(tag=TAG).error(f"实例化组件失败: {e}") def _init_prompt_enhancement(self): + # 更新上下文信息 self.prompt_manager.update_context_info(self, self.client_ip) enhanced_prompt = self.prompt_manager.build_enhanced_prompt( @@ -474,7 +515,11 @@ class ConnectionHandler: def _initialize_asr(self): """初始化ASR""" - if self._asr.interface_type == InterfaceType.LOCAL: + if ( + self._asr is not None + and hasattr(self._asr, "interface_type") + and self._asr.interface_type == InterfaceType.LOCAL + ): # 如果公共ASR是本地服务,则直接返回 # 因为本地一个实例ASR,可以被多个连接共享 asr = self._asr @@ -501,22 +546,35 @@ class ConnectionHandler: except Exception as e: self.logger.bind(tag=TAG).warning(f"声纹识别初始化失败: {str(e)}") - def _initialize_private_config(self): - """如果是从配置文件获取,则进行二次实例化""" + async def _background_initialize(self): + """在后台初始化配置和组件(完全不阻塞主循环)""" + try: + # 异步获取差异化配置 + await self._initialize_private_config_async() + # 在线程池中初始化组件 + self.executor.submit(self._initialize_components) + except Exception as e: + self.logger.bind(tag=TAG).error(f"后台初始化失败: {e}") + + async def _initialize_private_config_async(self): + """从接口异步获取差异化配置(异步版本,不阻塞主循环)""" if not self.read_config_from_api: + self.need_bind = False + self.bind_completed_event.set() return - """从接口获取差异化的配置进行二次实例化,非全量重新实例化""" try: begin_time = time.time() - private_config = get_private_config_from_api( + private_config = await get_private_config_from_api( self.config, self.headers.get("device-id"), self.headers.get("client-id", self.headers.get("device-id")), ) private_config["delete_audio"] = bool(self.config.get("delete_audio", True)) self.logger.bind(tag=TAG).info( - f"{time.time() - begin_time} 秒,获取差异化配置成功: {json.dumps(filter_sensitive_info(private_config), ensure_ascii=False)}" + f"{time.time() - begin_time} 秒,异步获取差异化配置成功: {json.dumps(filter_sensitive_info(private_config), ensure_ascii=False)}" ) + self.need_bind = False + self.bind_completed_event.set() except DeviceNotFoundException as e: self.need_bind = True private_config = {} @@ -526,7 +584,7 @@ class ConnectionHandler: private_config = {} except Exception as e: self.need_bind = True - self.logger.bind(tag=TAG).error(f"获取差异化配置失败: {e}") + self.logger.bind(tag=TAG).error(f"异步获取差异化配置失败: {e}") private_config = {} init_llm, init_tts, init_memory, init_intent = ( @@ -599,8 +657,14 @@ class ConnectionHandler: self.chat_history_conf = int(private_config["chat_history_conf"]) if private_config.get("mcp_endpoint", None) is not None: self.config["mcp_endpoint"] = private_config["mcp_endpoint"] + if private_config.get("context_providers", None) is not None: + self.config["context_providers"] = private_config["context_providers"] + + # 使用 run_in_executor 在线程池中执行 initialize_modules,避免阻塞主循环 try: - modules = initialize_modules( + modules = await self.loop.run_in_executor( + None, # 使用默认线程池 + initialize_modules, self.logger, private_config, init_vad, @@ -744,18 +808,26 @@ class ConnectionHandler: force_final_answer = False # 标记是否强制最终回答 if depth >= MAX_DEPTH: - self.logger.bind(tag=TAG).debug(f"已达到最大工具调用深度 {MAX_DEPTH},将强制基于现有信息回答") + self.logger.bind(tag=TAG).debug( + f"已达到最大工具调用深度 {MAX_DEPTH},将强制基于现有信息回答" + ) force_final_answer = True # 添加系统指令,要求 LLM 基于现有信息回答 - self.dialogue.put(Message( - role="user", - content="[系统提示] 已达到最大工具调用次数限制,请你基于目前已经获取的所有信息,直接给出最终答案。不要再尝试调用任何工具。" - )) + self.dialogue.put( + Message( + role="user", + content="[系统提示] 已达到最大工具调用次数限制,请你基于目前已经获取的所有信息,直接给出最终答案。不要再尝试调用任何工具。", + ) + ) # Define intent functions functions = None # 达到最大深度时,禁用工具调用,强制 LLM 直接回答 - if self.intent_type == "function_call" and hasattr(self, "func_handler") and not force_final_answer: + if ( + self.intent_type == "function_call" + and hasattr(self, "func_handler") + and not force_final_answer + ): functions = self.func_handler.get_functions() response_message = [] @@ -844,11 +916,16 @@ class ConnectionHandler: if a is not None: try: content_arguments_json = json.loads(a) - tool_calls_list.append({ - "id": str(uuid.uuid4().hex), - "name": content_arguments_json["name"], - "arguments": json.dumps(content_arguments_json["arguments"], ensure_ascii=False) - }) + tool_calls_list.append( + { + "id": str(uuid.uuid4().hex), + "name": content_arguments_json["name"], + "arguments": json.dumps( + content_arguments_json["arguments"], + ensure_ascii=False, + ), + } + ) except Exception as e: bHasError = True response_message.append(a) @@ -880,7 +957,9 @@ class ConnectionHandler: ) future = asyncio.run_coroutine_threadsafe( - self.func_handler.handle_llm_function_call(self, tool_call_data), + self.func_handler.handle_llm_function_call( + self, tool_call_data + ), self.loop, ) futures_with_data.append((future, tool_call_data)) @@ -888,7 +967,7 @@ class ConnectionHandler: # 等待协程结束(实际等待时长为最慢的那个) tool_results = [] for future, tool_call_data in futures_with_data: - result = future.result() + result = future.result() tool_results.append((result, tool_call_data)) # 统一处理所有工具调用结果 @@ -922,7 +1001,11 @@ class ConnectionHandler: need_llm_tools = [] for result, tool_call_data in tool_results: - if result.action in [Action.RESPONSE, Action.NOTFOUND, Action.ERROR]: # 直接回复前端 + if result.action in [ + Action.RESPONSE, + Action.NOTFOUND, + Action.ERROR, + ]: # 直接回复前端 text = result.response if result.response else result.result self.tts.tts_one_sentence(self, ContentType.TEXT, content_detail=text) self.dialogue.put(Message(role="assistant", content=text)) @@ -957,7 +1040,11 @@ class ConnectionHandler: self.dialogue.put( Message( role="tool", - tool_call_id=str(uuid.uuid4()) if tool_call_data["id"] is None else tool_call_data["id"], + tool_call_id=( + str(uuid.uuid4()) + if tool_call_data["id"] is None + else tool_call_data["id"] + ), content=text, ) ) @@ -990,8 +1077,8 @@ class ConnectionHandler: def _process_report(self, type, text, audio_data, report_time): """处理上报任务""" try: - # 执行上报(传入二进制数据) - report(self, type, text, audio_data, report_time) + # 执行异步上报(在事件循环中运行) + asyncio.run(report(self, type, text, audio_data, report_time)) except Exception as e: self.logger.bind(tag=TAG).error(f"上报处理异常: {e}") finally: @@ -1082,7 +1169,6 @@ class ConnectionHandler: f"关闭线程池时出错: {executor_error}" ) self.executor = None - self.logger.bind(tag=TAG).info("连接资源已释放") except Exception as e: self.logger.bind(tag=TAG).error(f"关闭连接时出错: {e}") @@ -1112,6 +1198,11 @@ class ConnectionHandler: except queue.Empty: break + # 重置音频流控器(取消后台任务并清空队列) + if hasattr(self, "audio_rate_controller") and self.audio_rate_controller: + self.audio_rate_controller.reset() + self.logger.bind(tag=TAG).debug("已重置音频流控器") + self.logger.bind(tag=TAG).debug( f"清理结束: TTS队列大小={self.tts.tts_text_queue.qsize()}, 音频队列大小={self.tts.tts_audio_queue.qsize()}" ) @@ -1137,13 +1228,14 @@ class ConnectionHandler: """检查连接超时""" try: while not self.stop_event.is_set(): + last_activity_time = self.last_activity_time + if self.need_bind: + last_activity_time = self.first_activity_time + # 检查是否超时(只有在时间戳已初始化的情况下) - if self.last_activity_time > 0.0: + if last_activity_time > 0.0: current_time = time.time() * 1000 - if ( - current_time - self.last_activity_time - > self.timeout_seconds * 1000 - ): + if current_time - last_activity_time > self.timeout_seconds * 1000: if not self.stop_event.is_set(): self.logger.bind(tag=TAG).info("连接超时,准备关闭") # 设置停止事件,防止重复处理 @@ -1171,7 +1263,7 @@ class ConnectionHandler: tools_call: 新的工具调用 """ for tool_call in tools_call: - tool_index = getattr(tool_call, 'index', None) + tool_index = getattr(tool_call, "index", None) if tool_index is None: if tool_call.function.name: # 有 function_name,说明是新的工具调用 @@ -1189,4 +1281,4 @@ class ConnectionHandler: if tool_call.function.name: tool_calls_list[tool_index]["name"] = tool_call.function.name if tool_call.function.arguments: - tool_calls_list[tool_index]["arguments"] += tool_call.function.arguments \ No newline at end of file + tool_calls_list[tool_index]["arguments"] += tool_call.function.arguments diff --git a/main/xiaozhi-server/core/handle/helloHandle.py b/main/xiaozhi-server/core/handle/helloHandle.py index 248c5f73..a4220d5a 100644 --- a/main/xiaozhi-server/core/handle/helloHandle.py +++ b/main/xiaozhi-server/core/handle/helloHandle.py @@ -101,7 +101,7 @@ async def checkWakeupWords(conn, text): } # 获取音频数据 - opus_packets = audio_to_data(response.get("file_path")) + opus_packets = await audio_to_data(response.get("file_path"), use_cache=False) # 播放唤醒词回复 conn.client_abort = False diff --git a/main/xiaozhi-server/core/handle/receiveAudioHandle.py b/main/xiaozhi-server/core/handle/receiveAudioHandle.py index 4eaf9ab1..8879564f 100644 --- a/main/xiaozhi-server/core/handle/receiveAudioHandle.py +++ b/main/xiaozhi-server/core/handle/receiveAudioHandle.py @@ -123,7 +123,7 @@ async def max_out_size(conn): text = "不好意思,我现在有点事情要忙,明天这个时候我们再聊,约好了哦!明天不见不散,拜拜!" await send_stt_message(conn, text) file_path = "config/assets/max_output_size.wav" - opus_packets = audio_to_data(file_path) + opus_packets = await audio_to_data(file_path) conn.tts.tts_audio_queue.put((SentenceType.LAST, opus_packets, text)) conn.close_after_chat = True @@ -142,7 +142,7 @@ async def check_bind_device(conn): # 播放提示音 music_path = "config/assets/bind_code.wav" - opus_packets = audio_to_data(music_path) + opus_packets = await audio_to_data(music_path) conn.tts.tts_audio_queue.put((SentenceType.FIRST, opus_packets, text)) # 逐个播放数字 @@ -150,7 +150,7 @@ async def check_bind_device(conn): try: digit = conn.bind_code[i] num_path = f"config/assets/bind_code/{digit}.wav" - num_packets = audio_to_data(num_path) + num_packets = await audio_to_data(num_path) conn.tts.tts_audio_queue.put((SentenceType.MIDDLE, num_packets, None)) except Exception as e: conn.logger.bind(tag=TAG).error(f"播放数字音频失败: {e}") @@ -162,5 +162,5 @@ async def check_bind_device(conn): text = f"没有找到该设备的版本信息,请正确配置 OTA地址,然后重新编译固件。" await send_stt_message(conn, text) music_path = "config/assets/bind_not_found.wav" - opus_packets = audio_to_data(music_path) + opus_packets = await audio_to_data(music_path) conn.tts.tts_audio_queue.put((SentenceType.LAST, opus_packets, text)) diff --git a/main/xiaozhi-server/core/handle/reportHandle.py b/main/xiaozhi-server/core/handle/reportHandle.py index 7b30f79c..053e8f2e 100644 --- a/main/xiaozhi-server/core/handle/reportHandle.py +++ b/main/xiaozhi-server/core/handle/reportHandle.py @@ -10,7 +10,6 @@ TTS上报功能已集成到ConnectionHandler类中。 """ import time - import opuslib_next from config.manage_api_client import report as manage_report @@ -18,7 +17,7 @@ from config.manage_api_client import report as manage_report TAG = __name__ -def report(conn, type, text, opus_data, report_time): +async def report(conn, type, text, opus_data, report_time): """执行聊天记录上报操作 Args: @@ -33,8 +32,8 @@ def report(conn, type, text, opus_data, report_time): audio_data = opus_to_wav(conn, opus_data) else: audio_data = None - # 执行上报 - manage_report( + # 执行异步上报 + await manage_report( mac_address=conn.device_id, session_id=conn.session_id, chat_type=type, @@ -56,41 +55,49 @@ def opus_to_wav(conn, opus_data): Returns: bytes: WAV格式的音频数据 """ - decoder = opuslib_next.Decoder(16000, 1) # 16kHz, 单声道 - pcm_data = [] + decoder = None + try: + decoder = opuslib_next.Decoder(16000, 1) # 16kHz, 单声道 + pcm_data = [] - for opus_packet in opus_data: - try: - pcm_frame = decoder.decode(opus_packet, 960) # 960 samples = 60ms - pcm_data.append(pcm_frame) - except opuslib_next.OpusError as e: - conn.logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True) + for opus_packet in opus_data: + try: + pcm_frame = decoder.decode(opus_packet, 960) # 960 samples = 60ms + pcm_data.append(pcm_frame) + except opuslib_next.OpusError as e: + conn.logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True) - if not pcm_data: - raise ValueError("没有有效的PCM数据") + if not pcm_data: + raise ValueError("没有有效的PCM数据") - # 创建WAV文件头 - pcm_data_bytes = b"".join(pcm_data) - num_samples = len(pcm_data_bytes) // 2 # 16-bit samples + # 创建WAV文件头 + pcm_data_bytes = b"".join(pcm_data) + num_samples = len(pcm_data_bytes) // 2 # 16-bit samples - # WAV文件头 - wav_header = bytearray() - wav_header.extend(b"RIFF") # ChunkID - wav_header.extend((36 + len(pcm_data_bytes)).to_bytes(4, "little")) # ChunkSize - wav_header.extend(b"WAVE") # Format - wav_header.extend(b"fmt ") # Subchunk1ID - wav_header.extend((16).to_bytes(4, "little")) # Subchunk1Size - wav_header.extend((1).to_bytes(2, "little")) # AudioFormat (PCM) - wav_header.extend((1).to_bytes(2, "little")) # NumChannels - wav_header.extend((16000).to_bytes(4, "little")) # SampleRate - wav_header.extend((32000).to_bytes(4, "little")) # ByteRate - wav_header.extend((2).to_bytes(2, "little")) # BlockAlign - wav_header.extend((16).to_bytes(2, "little")) # BitsPerSample - wav_header.extend(b"data") # Subchunk2ID - wav_header.extend(len(pcm_data_bytes).to_bytes(4, "little")) # Subchunk2Size + # WAV文件头 + wav_header = bytearray() + wav_header.extend(b"RIFF") # ChunkID + wav_header.extend((36 + len(pcm_data_bytes)).to_bytes(4, "little")) # ChunkSize + wav_header.extend(b"WAVE") # Format + wav_header.extend(b"fmt ") # Subchunk1ID + wav_header.extend((16).to_bytes(4, "little")) # Subchunk1Size + wav_header.extend((1).to_bytes(2, "little")) # AudioFormat (PCM) + wav_header.extend((1).to_bytes(2, "little")) # NumChannels + wav_header.extend((16000).to_bytes(4, "little")) # SampleRate + wav_header.extend((32000).to_bytes(4, "little")) # ByteRate + wav_header.extend((2).to_bytes(2, "little")) # BlockAlign + wav_header.extend((16).to_bytes(2, "little")) # BitsPerSample + wav_header.extend(b"data") # Subchunk2ID + wav_header.extend(len(pcm_data_bytes).to_bytes(4, "little")) # Subchunk2Size - # 返回完整的WAV数据 - return bytes(wav_header) + pcm_data_bytes + # 返回完整的WAV数据 + return bytes(wav_header) + pcm_data_bytes + finally: + if decoder is not None: + try: + del decoder + except Exception as e: + conn.logger.bind(tag=TAG).debug(f"释放decoder资源时出错: {e}") def enqueue_tts_report(conn, text, opus_data): diff --git a/main/xiaozhi-server/core/handle/sendAudioHandle.py b/main/xiaozhi-server/core/handle/sendAudioHandle.py index d517cdb9..8198dfb0 100644 --- a/main/xiaozhi-server/core/handle/sendAudioHandle.py +++ b/main/xiaozhi-server/core/handle/sendAudioHandle.py @@ -4,8 +4,13 @@ import asyncio from core.utils import textUtils from core.utils.util import audio_to_data from core.providers.tts.dto.dto import SentenceType +from core.utils.audioRateController import AudioRateController TAG = __name__ +# 音频帧时长(毫秒) +AUDIO_FRAME_DURATION = 60 +# 预缓冲包数量,直接发送以减少延迟 +PRE_BUFFER_COUNT = 5 async def sendAudioMessage(conn, sentenceType, audios, text): @@ -15,7 +20,19 @@ async def sendAudioMessage(conn, sentenceType, audios, text): await send_tts_message(conn, "start", None) if sentenceType == SentenceType.FIRST: - await send_tts_message(conn, "sentence_start", text) + # 同一句子的后续消息加入流控队列,其他情况立即发送 + if ( + hasattr(conn, "audio_rate_controller") + and conn.audio_rate_controller + and getattr(conn, "audio_flow_control", {}).get("sentence_id") + == conn.sentence_id + ): + conn.audio_rate_controller.add_message( + lambda: send_tts_message(conn, "sentence_start", text) + ) + else: + # 新句子或流控器未初始化,立即发送 + await send_tts_message(conn, "sentence_start", text) await sendAudio(conn, audios) # 发送句子开始消息 @@ -30,29 +47,27 @@ async def sendAudioMessage(conn, sentenceType, audios, text): await conn.close() -def calculate_timestamp_and_sequence(conn, start_time, packet_index, frame_duration=60): +async def _wait_for_audio_completion(conn): """ - 计算音频数据包的时间戳和序列号 + 等待音频队列清空并等待预缓冲包播放完成 + Args: conn: 连接对象 - start_time: 起始时间(性能计数器值) - packet_index: 数据包索引 - frame_duration: 帧时长(毫秒),匹配 Opus 编码 - Returns: - tuple: (timestamp, sequence) """ - # 计算时间戳(使用播放位置计算) - timestamp = int((start_time + packet_index * frame_duration / 1000) * 1000) % ( - 2**32 - ) + if hasattr(conn, "audio_rate_controller") and conn.audio_rate_controller: + rate_controller = conn.audio_rate_controller + conn.logger.bind(tag=TAG).debug( + f"等待音频发送完成,队列中还有 {len(rate_controller.queue)} 个包" + ) + await rate_controller.queue_empty_event.wait() - # 计算序列号 - if hasattr(conn, "audio_flow_control"): - sequence = conn.audio_flow_control["sequence"] - else: - sequence = packet_index # 如果没有流控状态,直接使用索引 + # 等待预缓冲包播放完成 + # 前N个包直接发送,增加2个网络抖动包,需要额外等待它们在客户端播放完成 + frame_duration_ms = rate_controller.frame_duration + pre_buffer_playback_time = (PRE_BUFFER_COUNT + 2) * frame_duration_ms / 1000.0 + await asyncio.sleep(pre_buffer_playback_time) - return timestamp, sequence + conn.logger.bind(tag=TAG).debug("音频发送完成") async def _send_to_mqtt_gateway(conn, opus_packet, timestamp, sequence): @@ -77,135 +92,151 @@ async def _send_to_mqtt_gateway(conn, opus_packet, timestamp, sequence): await conn.websocket.send(complete_packet) -# 播放音频 -async def sendAudio(conn, audios, frame_duration=60): +async def sendAudio(conn, audios, frame_duration=AUDIO_FRAME_DURATION): """ - 发送单个opus包,支持流控 + 发送音频包,使用 AudioRateController 进行精确的流量控制 + Args: conn: 连接对象 - opus_packet: 单个opus数据包 - pre_buffer: 快速发送音频 - frame_duration: 帧时长(毫秒),匹配 Opus 编码 + audios: 单个opus包(bytes) 或 opus包列表 + frame_duration: 帧时长(毫秒),默认使用全局常量AUDIO_FRAME_DURATION """ if audios is None or len(audios) == 0: return - # 获取发送延迟配置 send_delay = conn.config.get("tts_audio_send_delay", -1) / 1000.0 + is_single_packet = isinstance(audios, bytes) - if isinstance(audios, bytes): - # 重置流控状态,第一次读取和会话发生转变时 - if not hasattr(conn, "audio_flow_control") or conn.audio_flow_control.get("sentence_id") != conn.sentence_id: - conn.audio_flow_control = { - "last_send_time": 0, - "packet_count": 0, - "start_time": time.perf_counter(), - "sequence": 0, # 添加序列号 - "sentence_id": conn.sentence_id, - } + # 初始化或获取 RateController + rate_controller, flow_control = _get_or_create_rate_controller( + conn, frame_duration, is_single_packet + ) + # 统一转换为列表处理 + audio_list = [audios] if is_single_packet else audios + + # 发送音频包 + await _send_audio_with_rate_control( + conn, audio_list, rate_controller, flow_control, send_delay + ) + + +def _get_or_create_rate_controller(conn, frame_duration, is_single_packet): + """ + 获取或创建 RateController 和 flow_control + + Args: + conn: 连接对象 + frame_duration: 帧时长 + is_single_packet: 是否单包模式(True: TTS流式单包, False: 批量包) + + Returns: + (rate_controller, flow_control) + """ + # 判断是否需要重置:单包模式且 sentence_id 变化,或者控制器不存在 + need_reset = ( + is_single_packet + and getattr(conn, "audio_flow_control", {}).get("sentence_id") + != conn.sentence_id + ) or not hasattr(conn, "audio_rate_controller") + + if need_reset: + # 创建或获取 rate_controller + if not hasattr(conn, "audio_rate_controller"): + conn.audio_rate_controller = AudioRateController(frame_duration) + else: + conn.audio_rate_controller.reset() + + # 初始化 flow_control + conn.audio_flow_control = { + "packet_count": 0, + "sequence": 0, + "sentence_id": conn.sentence_id, + } + + # 启动后台发送循环 + _start_background_sender( + conn, conn.audio_rate_controller, conn.audio_flow_control + ) + + return conn.audio_rate_controller, conn.audio_flow_control + + +def _start_background_sender(conn, rate_controller, flow_control): + """ + 启动后台发送循环任务 + + Args: + conn: 连接对象 + rate_controller: 速率控制器 + flow_control: 流控状态 + """ + + async def send_callback(packet): + # 检查是否应该中止 + if conn.client_abort: + raise asyncio.CancelledError("客户端已中止") + + conn.last_activity_time = time.time() * 1000 + await _do_send_audio(conn, packet, flow_control) + conn.client_is_speaking = True + + # 使用 start_sending 启动后台循环 + rate_controller.start_sending(send_callback) + + +async def _send_audio_with_rate_control( + conn, audio_list, rate_controller, flow_control, send_delay +): + """ + 使用 rate_controller 发送音频包 + + Args: + conn: 连接对象 + audio_list: 音频包列表 + rate_controller: 速率控制器 + flow_control: 流控状态 + send_delay: 固定延迟(秒),-1表示使用动态流控 + """ + for packet in audio_list: if conn.client_abort: return conn.last_activity_time = time.time() * 1000 - # 预缓冲:前5个包直接发送,不做延迟 - pre_buffer_count = 5 - flow_control = conn.audio_flow_control - current_time = time.perf_counter() - - if flow_control["packet_count"] < pre_buffer_count: - # 预缓冲阶段,直接发送不延迟 - pass + # 预缓冲:前N个包直接发送 + if flow_control["packet_count"] < PRE_BUFFER_COUNT: + await _do_send_audio(conn, packet, flow_control) + conn.client_is_speaking = True elif send_delay > 0: - # 使用固定延迟 + # 固定延迟模式 await asyncio.sleep(send_delay) + await _do_send_audio(conn, packet, flow_control) + conn.client_is_speaking = True else: - effective_packet = flow_control["packet_count"] - pre_buffer_count - expected_time = flow_control["start_time"] + ( - effective_packet * frame_duration / 1000 - ) - delay = expected_time - current_time - if delay > 0: - await asyncio.sleep(delay) - else: - # 纠正误差 - flow_control["start_time"] += abs(delay) + # 动态流控模式:仅添加到队列,由后台循环负责发送 + rate_controller.add_audio(packet) - if conn.conn_from_mqtt_gateway: - # 计算时间戳和序列号 - timestamp, sequence = calculate_timestamp_and_sequence( - conn, - flow_control["start_time"], - flow_control["packet_count"], - frame_duration, - ) - # 调用通用函数发送带头部的数据包 - await _send_to_mqtt_gateway(conn, audios, timestamp, sequence) - else: - # 直接发送opus数据包,不添加头部 - await conn.websocket.send(audios) - conn.client_is_speaking = True - # 更新流控状态 - flow_control["packet_count"] += 1 - flow_control["sequence"] += 1 - flow_control["last_send_time"] = time.perf_counter() +async def _do_send_audio(conn, opus_packet, flow_control): + """ + 执行实际的音频发送 + """ + packet_index = flow_control.get("packet_count", 0) + sequence = flow_control.get("sequence", 0) + + if conn.conn_from_mqtt_gateway: + # 计算时间戳(基于播放位置) + start_time = time.time() + timestamp = int(start_time * 1000) % (2**32) + await _send_to_mqtt_gateway(conn, opus_packet, timestamp, sequence) else: - # 文件型音频走普通播放 - start_time = time.perf_counter() - play_position = 0 + # 直接发送opus数据包 + await conn.websocket.send(opus_packet) - # 执行预缓冲 - pre_buffer_frames = min(5, len(audios)) - for i in range(pre_buffer_frames): - if conn.conn_from_mqtt_gateway: - # 计算时间戳和序列号 - timestamp, sequence = calculate_timestamp_and_sequence( - conn, start_time, i, frame_duration - ) - # 调用通用函数发送带头部的数据包 - await _send_to_mqtt_gateway(conn, audios[i], timestamp, sequence) - else: - # 直接发送预缓冲包,不添加头部 - await conn.websocket.send(audios[i]) - conn.client_is_speaking = True - remaining_audios = audios[pre_buffer_frames:] - - # 播放剩余音频帧 - for i, opus_packet in enumerate(remaining_audios): - if conn.client_abort: - break - - # 重置没有声音的状态 - conn.last_activity_time = time.time() * 1000 - - if send_delay > 0: - # 固定延迟模式 - await asyncio.sleep(send_delay) - else: - # 计算预期发送时间 - expected_time = start_time + (play_position / 1000) - current_time = time.perf_counter() - delay = expected_time - current_time - if delay > 0: - await asyncio.sleep(delay) - - if conn.conn_from_mqtt_gateway: - # 计算时间戳和序列号(使用当前的数据包索引确保连续性) - packet_index = pre_buffer_frames + i - timestamp, sequence = calculate_timestamp_and_sequence( - conn, start_time, packet_index, frame_duration - ) - # 调用通用函数发送带头部的数据包 - await _send_to_mqtt_gateway(conn, opus_packet, timestamp, sequence) - else: - # 直接发送opus数据包,不添加头部 - await conn.websocket.send(opus_packet) - - conn.client_is_speaking = True - - play_position += frame_duration + # 更新流控状态 + flow_control["packet_count"] = packet_index + 1 + flow_control["sequence"] = sequence + 1 async def send_tts_message(conn, state, text=None): @@ -224,8 +255,10 @@ async def send_tts_message(conn, state, text=None): stop_tts_notify_voice = conn.config.get( "stop_tts_notify_voice", "config/assets/tts_notify.mp3" ) - audios = audio_to_data(stop_tts_notify_voice, is_opus=True) + audios = await audio_to_data(stop_tts_notify_voice, is_opus=True) await sendAudio(conn, audios) + # 等待所有音频包发送完成 + await _wait_for_audio_completion(conn) # 清除服务端讲话状态 conn.clearSpeakStatus() diff --git a/main/xiaozhi-server/core/handle/textHandler/listenMessageHandler.py b/main/xiaozhi-server/core/handle/textHandler/listenMessageHandler.py index 97286dfe..c71649ed 100644 --- a/main/xiaozhi-server/core/handle/textHandler/listenMessageHandler.py +++ b/main/xiaozhi-server/core/handle/textHandler/listenMessageHandler.py @@ -1,12 +1,14 @@ import time +import asyncio from typing import Dict, Any -from core.handle.receiveAudioHandle import handleAudioMessage, startToChat +from core.handle.receiveAudioHandle import startToChat from core.handle.reportHandle import enqueue_asr_report from core.handle.sendAudioHandle import send_stt_message, send_tts_message from core.handle.textMessageHandler import TextMessageHandler from core.handle.textMessageType import TextMessageType from core.utils.util import remove_punctuation_and_length +from core.providers.asr.dto.dto import InterfaceType TAG = __name__ @@ -29,8 +31,18 @@ class ListenTextMessageHandler(TextMessageHandler): 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"") + if conn.asr.interface_type == InterfaceType.STREAM: + # 流式模式下,发送结束请求 + asyncio.create_task(conn.asr._send_stop_request()) + else: + # 非流式模式:直接触发ASR识别 + if len(conn.asr_audio) > 0: + asr_audio_task = conn.asr_audio.copy() + conn.asr_audio.clear() + conn.reset_vad_states() + + if len(asr_audio_task) > 0: + await conn.asr.handle_voice_stop(conn, asr_audio_task) elif msg_json["state"] == "detect": conn.client_have_voice = False conn.asr_audio.clear() @@ -57,6 +69,7 @@ class ListenTextMessageHandler(TextMessageHandler): enqueue_asr_report(conn, "嘿,你好呀", []) await startToChat(conn, "嘿,你好呀") else: + conn.just_woken_up = True # 上报纯文字数据(复用ASR上报功能,但不提供音频数据) enqueue_asr_report(conn, original_text, []) # 否则需要LLM对文字内容进行答复 diff --git a/main/xiaozhi-server/core/handle/textHandler/pingMessageHandler.py b/main/xiaozhi-server/core/handle/textHandler/pingMessageHandler.py new file mode 100644 index 00000000..f3af85eb --- /dev/null +++ b/main/xiaozhi-server/core/handle/textHandler/pingMessageHandler.py @@ -0,0 +1,45 @@ +import json +import time +from typing import Dict, Any + +from core.handle.textMessageHandler import TextMessageHandler +from core.handle.textMessageType import TextMessageType + +TAG = __name__ + + +class PingMessageHandler(TextMessageHandler): + """Ping消息处理器,用于保持WebSocket连接""" + + @property + def message_type(self) -> TextMessageType: + return TextMessageType.PING + + async def handle(self, conn, msg_json: Dict[str, Any]) -> None: + """ + 处理PING消息,发送PONG响应 + 消息格式:{"type": "ping"} + Args: + conn: WebSocket连接对象 + msg_json: PING消息的JSON数据 + """ + # 检查是否启用了WebSocket心跳功能 + enable_websocket_ping = conn.config.get("enable_websocket_ping", False) + if not enable_websocket_ping: + conn.logger.debug(f"WebSocket心跳功能未启用,忽略PING消息") + return + + try: + conn.logger.debug(f"收到PING消息,发送PONG响应") + conn.last_activity_time = time.time() * 1000 + # 构造PONG响应消息 + pong_message = { + "type": "pong", + "timestamp": time.strftime("%Y-%m-%d %H:%M:%S", time.localtime()), + } + + # 发送PONG响应 + await conn.websocket.send(json.dumps(pong_message)) + + except Exception as e: + conn.logger.error(f"处理PING消息时发生错误: {e}") diff --git a/main/xiaozhi-server/core/handle/textMessageHandlerRegistry.py b/main/xiaozhi-server/core/handle/textMessageHandlerRegistry.py index e90d7231..65e9474f 100644 --- a/main/xiaozhi-server/core/handle/textMessageHandlerRegistry.py +++ b/main/xiaozhi-server/core/handle/textMessageHandlerRegistry.py @@ -7,6 +7,7 @@ from core.handle.textHandler.listenMessageHandler import ListenTextMessageHandle from core.handle.textHandler.mcpMessageHandler import McpTextMessageHandler from core.handle.textMessageHandler import TextMessageHandler from core.handle.textHandler.serverMessageHandler import ServerTextMessageHandler +from core.handle.textHandler.pingMessageHandler import PingMessageHandler TAG = __name__ @@ -27,6 +28,7 @@ class TextMessageHandlerRegistry: IotTextMessageHandler(), McpTextMessageHandler(), ServerTextMessageHandler(), + PingMessageHandler(), ] for handler in handlers: diff --git a/main/xiaozhi-server/core/handle/textMessageType.py b/main/xiaozhi-server/core/handle/textMessageType.py index 53e71b71..bd04d289 100644 --- a/main/xiaozhi-server/core/handle/textMessageType.py +++ b/main/xiaozhi-server/core/handle/textMessageType.py @@ -9,3 +9,4 @@ class TextMessageType(Enum): IOT = "iot" MCP = "mcp" SERVER = "server" + PING = "ping" diff --git a/main/xiaozhi-server/core/http_server.py b/main/xiaozhi-server/core/http_server.py index edbdf1fe..feb96f3b 100644 --- a/main/xiaozhi-server/core/http_server.py +++ b/main/xiaozhi-server/core/http_server.py @@ -33,38 +33,60 @@ class SimpleHttpServer: return f"ws://{local_ip}:{port}/xiaozhi/v1/" async def start(self): - server_config = self.config["server"] - read_config_from_api = self.config.get("read_config_from_api", False) - host = server_config.get("ip", "0.0.0.0") - port = int(server_config.get("http_port", 8003)) + try: + server_config = self.config["server"] + read_config_from_api = self.config.get("read_config_from_api", False) + host = server_config.get("ip", "0.0.0.0") + port = int(server_config.get("http_port", 8003)) - if port: - app = web.Application() + if port: + app = web.Application() - if not read_config_from_api: - # 如果没有开启智控台,只是单模块运行,就需要再添加简单OTA接口,用于下发websocket接口 + if not read_config_from_api: + # 如果没有开启智控台,只是单模块运行,就需要再添加简单OTA接口,用于下发websocket接口 + app.add_routes( + [ + web.get("/xiaozhi/ota/", self.ota_handler.handle_get), + web.post("/xiaozhi/ota/", self.ota_handler.handle_post), + web.options( + "/xiaozhi/ota/", self.ota_handler.handle_options + ), + # 下载接口,仅提供 data/bin/*.bin 下载 + web.get( + "/xiaozhi/ota/download/{filename}", + self.ota_handler.handle_download, + ), + web.options( + "/xiaozhi/ota/download/{filename}", + self.ota_handler.handle_options, + ), + ] + ) + # 添加路由 app.add_routes( [ - web.get("/xiaozhi/ota/", self.ota_handler.handle_get), - web.post("/xiaozhi/ota/", self.ota_handler.handle_post), - web.options("/xiaozhi/ota/", self.ota_handler.handle_post), + web.get("/mcp/vision/explain", self.vision_handler.handle_get), + web.post( + "/mcp/vision/explain", self.vision_handler.handle_post + ), + web.options( + "/mcp/vision/explain", self.vision_handler.handle_options + ), ] ) - # 添加路由 - app.add_routes( - [ - web.get("/mcp/vision/explain", self.vision_handler.handle_get), - web.post("/mcp/vision/explain", self.vision_handler.handle_post), - web.options("/mcp/vision/explain", self.vision_handler.handle_post), - ] - ) - # 运行服务 - runner = web.AppRunner(app) - await runner.setup() - site = web.TCPSite(runner, host, port) - await site.start() + # 运行服务 + runner = web.AppRunner(app) + await runner.setup() + site = web.TCPSite(runner, host, port) + await site.start() - # 保持服务运行 - while True: - await asyncio.sleep(3600) # 每隔 1 小时检查一次 + # 保持服务运行 + while True: + await asyncio.sleep(3600) # 每隔 1 小时检查一次 + except Exception as e: + self.logger.bind(tag=TAG).error(f"HTTP服务器启动失败: {e}") + import traceback + + self.logger.bind(tag=TAG).error(f"错误堆栈: {traceback.format_exc()}") + raise diff --git a/main/xiaozhi-server/core/providers/asr/aliyun_stream.py b/main/xiaozhi-server/core/providers/asr/aliyun_stream.py index 8acc640a..4ce588a7 100644 --- a/main/xiaozhi-server/core/providers/asr/aliyun_stream.py +++ b/main/xiaozhi-server/core/providers/asr/aliyun_stream.py @@ -8,8 +8,6 @@ import asyncio import requests import websockets import opuslib_next -import random -from typing import Optional, Tuple, List from urllib import parse from datetime import datetime from config.logger import setup_logging @@ -139,13 +137,13 @@ class ASRProvider(ASRProviderBase): conn.asr_audio.append(audio) conn.asr_audio = conn.asr_audio[-10:] - # 只在有声音且没有连接时建立连接 - if audio_have_voice and not self.is_processing: + # 只在有声音且没有连接时建立连接(排除正在停止的情况) + if audio_have_voice and not self.is_processing and not self.asr_ws: try: await self._start_recognition(conn) except Exception as e: logger.bind(tag=TAG).error(f"开始识别失败: {str(e)}") - await self._cleanup(conn) + await self._cleanup() return if self.asr_ws and self.is_processing and self.server_ready: @@ -185,10 +183,8 @@ class ASRProvider(ASRProviderBase): "header": { "namespace": "SpeechTranscriber", "name": "StartTranscription", - "status": 20000000, "message_id": uuid.uuid4().hex, "task_id": self.task_id, - "status_text": "Gateway:SUCCESS:Success.", "appkey": self.appkey }, "payload": { @@ -207,18 +203,21 @@ class ASRProvider(ASRProviderBase): async def _forward_results(self, conn): """转发识别结果""" try: - while self.asr_ws and not conn.stop_event.is_set(): + while not conn.stop_event.is_set(): try: response = await asyncio.wait_for(self.asr_ws.recv(), timeout=1.0) result = json.loads(response) - + header = result.get("header", {}) payload = result.get("payload", {}) message_name = header.get("name", "") status = header.get("status", 0) - + if status != 20000000: - if status in [40000004, 40010004]: # 连接超时或客户端断开 + if status == 40010004: + logger.bind(tag=TAG).warning(f"请在服务端响应完成后再关闭链接,状态码: {status}") + break + if status in [40000004, 40010003]: # 连接超时或客户端断开 logger.bind(tag=TAG).warning(f"连接问题,状态码: {status}") break elif status in [40270002, 40270003]: # 音频问题 @@ -227,12 +226,12 @@ class ASRProvider(ASRProviderBase): else: logger.bind(tag=TAG).error(f"识别错误,状态码: {status}, 消息: {header.get('status_text', '')}") continue - + # 收到TranscriptionStarted表示服务器准备好接收音频数据 if message_name == "TranscriptionStarted": self.server_ready = True logger.bind(tag=TAG).debug("服务器已准备,开始发送缓存音频...") - + # 发送缓存音频 if conn.asr_audio: for cached_audio in conn.asr_audio[-10:]: @@ -243,89 +242,89 @@ class ASRProvider(ASRProviderBase): logger.bind(tag=TAG).warning(f"发送缓存音频失败: {e}") break continue - - if message_name == "TranscriptionResultChanged": - # 中间结果 - text = payload.get("result", "") - if text: - self.text = text elif message_name == "SentenceEnd": - # 最终结果 + # 句子结束(每个句子都会触发) text = payload.get("result", "") if text: - self.text = text - conn.reset_vad_states() - # 传递缓存的音频数据 - audio_data = getattr(conn, 'asr_audio_for_voiceprint', []) - await self.handle_voice_stop(conn, audio_data) - # 清空缓存 - conn.asr_audio_for_voiceprint = [] - break - elif message_name == "TranscriptionCompleted": - # 识别完成 - self.is_processing = False - break - + logger.bind(tag=TAG).info(f"识别到文本: {text}") + + # 手动模式下累积识别结果 + if conn.client_listen_mode == "manual": + if self.text: + self.text += text + else: + self.text = text + + # 手动模式下,只有在收到stop信号后才触发处理(仅处理一次) + if conn.client_voice_stop: + audio_data = getattr(conn, 'asr_audio_for_voiceprint', []) + if len(audio_data) > 0: + logger.bind(tag=TAG).debug("收到最终识别结果,触发处理") + await self.handle_voice_stop(conn, audio_data) + # 清理音频缓存 + conn.asr_audio.clear() + conn.reset_vad_states() + break + else: + # 自动模式下直接覆盖 + self.text = text + conn.reset_vad_states() + audio_data = getattr(conn, 'asr_audio_for_voiceprint', []) + await self.handle_voice_stop(conn, audio_data) + break + except asyncio.TimeoutError: - continue - except websockets.exceptions.ConnectionClosed: + logger.bind(tag=TAG).error("接收结果超时") + break + except websockets.ConnectionClosed: + logger.bind(tag=TAG).info("ASR服务连接已关闭") + self.is_processing = False break except Exception as e: logger.bind(tag=TAG).error(f"处理结果失败: {str(e)}") break - + except Exception as e: logger.bind(tag=TAG).error(f"结果转发失败: {str(e)}") finally: - await self._cleanup(conn) + # 清理连接的音频缓存 + await self._cleanup() + if conn: + if hasattr(conn, 'asr_audio_for_voiceprint'): + conn.asr_audio_for_voiceprint = [] + if hasattr(conn, 'asr_audio'): + conn.asr_audio = [] - async def _cleanup(self, conn): - """清理资源""" - logger.bind(tag=TAG).debug(f"开始ASR会话清理 | 当前状态: processing={self.is_processing}, server_ready={self.server_ready}") - - # 清理连接的音频缓存 - if conn and hasattr(conn, 'asr_audio_for_voiceprint'): - conn.asr_audio_for_voiceprint = [] - - # 判断是否需要发送终止请求 - should_stop = self.is_processing or self.server_ready - - # 发送停止识别请求 - if self.asr_ws and should_stop: + async def _send_stop_request(self): + """发送停止识别请求(不关闭连接)""" + if self.asr_ws: try: + # 先停止音频发送 + self.is_processing = False + stop_msg = { "header": { "namespace": "SpeechTranscriber", "name": "StopTranscription", - "status": 20000000, "message_id": uuid.uuid4().hex, "task_id": self.task_id, - "status_text": "Client:Stop", "appkey": self.appkey } } - logger.bind(tag=TAG).debug("正在发送ASR终止请求") + logger.bind(tag=TAG).debug("停止识别请求已发送") await self.asr_ws.send(json.dumps(stop_msg, ensure_ascii=False)) - await asyncio.sleep(0.1) - logger.bind(tag=TAG).debug("ASR终止请求已发送") except Exception as e: - logger.bind(tag=TAG).error(f"ASR终止请求发送失败: {e}") - - # 状态重置(在终止请求发送后) + logger.bind(tag=TAG).error(f"发送停止识别请求失败: {e}") + + async def _cleanup(self): + """清理资源(关闭连接)""" + logger.bind(tag=TAG).debug(f"开始ASR会话清理 | 当前状态: processing={self.is_processing}, server_ready={self.server_ready}") + + # 状态重置 self.is_processing = False self.server_ready = False logger.bind(tag=TAG).debug("ASR状态已重置") - # 清理任务 - if self.forward_task and not self.forward_task.done(): - self.forward_task.cancel() - try: - await asyncio.wait_for(self.forward_task, timeout=1.0) - except Exception as e: - logger.bind(tag=TAG).debug(f"forward_task取消异常: {e}") - finally: - self.forward_task = None - # 关闭连接 if self.asr_ws: try: @@ -336,7 +335,10 @@ class ASRProvider(ASRProviderBase): logger.bind(tag=TAG).error(f"关闭WebSocket连接失败: {e}") finally: self.asr_ws = None - + + # 清理任务引用 + self.forward_task = None + logger.bind(tag=TAG).debug("ASR会话清理完成") async def speech_to_text(self, opus_data, session_id, audio_format): @@ -347,4 +349,11 @@ class ASRProvider(ASRProviderBase): async def close(self): """关闭资源""" - await self._cleanup() + await self._cleanup(None) + if hasattr(self, 'decoder') and self.decoder is not None: + try: + del self.decoder + self.decoder = None + logger.bind(tag=TAG).debug("Aliyun decoder resources released") + except Exception as e: + logger.bind(tag=TAG).debug(f"释放Aliyun decoder资源时出错: {e}") diff --git a/main/xiaozhi-server/core/providers/asr/base.py b/main/xiaozhi-server/core/providers/asr/base.py index ed69da89..6c6969fd 100644 --- a/main/xiaozhi-server/core/providers/asr/base.py +++ b/main/xiaozhi-server/core/providers/asr/base.py @@ -9,7 +9,6 @@ import asyncio import traceback import threading import opuslib_next -import concurrent.futures from abc import ABC, abstractmethod from config.logger import setup_logging from typing import Optional, Tuple, List @@ -53,121 +52,89 @@ class ASRProviderBase(ABC): # 接收音频 async def receive_audio(self, conn, audio, audio_have_voice): - if conn.client_listen_mode == "auto" or conn.client_listen_mode == "realtime": - have_voice = audio_have_voice + if conn.client_listen_mode == "manual": + # 手动模式:缓存音频用于ASR识别 + conn.asr_audio.append(audio) else: - have_voice = conn.client_have_voice - - conn.asr_audio.append(audio) - if not have_voice and not conn.client_have_voice: - conn.asr_audio = conn.asr_audio[-10:] - return + # 自动/实时模式:使用VAD检测 + have_voice = audio_have_voice - if conn.client_voice_stop: - asr_audio_task = conn.asr_audio.copy() - conn.asr_audio.clear() - conn.reset_vad_states() + conn.asr_audio.append(audio) + if not have_voice and not conn.client_have_voice: + conn.asr_audio = conn.asr_audio[-10:] + return - if len(asr_audio_task) > 15: - await self.handle_voice_stop(conn, asr_audio_task) + # 自动模式下通过VAD检测到语音停止时触发识别 + if conn.client_voice_stop: + asr_audio_task = conn.asr_audio.copy() + conn.asr_audio.clear() + conn.reset_vad_states() + + if len(asr_audio_task) > 15: + await self.handle_voice_stop(conn, asr_audio_task) # 处理语音停止 async def handle_voice_stop(self, conn, asr_audio_task: List[bytes]): """并行处理ASR和声纹识别""" try: total_start_time = time.monotonic() - + # 准备音频数据 if conn.audio_format == "pcm": pcm_data = asr_audio_task else: pcm_data = self.decode_opus(asr_audio_task) - + combined_pcm_data = b"".join(pcm_data) - + # 预先准备WAV数据 wav_data = None if conn.voiceprint_provider and combined_pcm_data: wav_data = self._pcm_to_wav(combined_pcm_data) - + # 定义ASR任务 - def run_asr(): - start_time = time.monotonic() - try: - loop = asyncio.new_event_loop() - asyncio.set_event_loop(loop) - try: - result = loop.run_until_complete( - self.speech_to_text(asr_audio_task, conn.session_id, conn.audio_format) - ) - end_time = time.monotonic() - logger.bind(tag=TAG).debug(f"ASR耗时: {end_time - start_time:.3f}s") - return result - finally: - loop.close() - except Exception as e: - end_time = time.monotonic() - logger.bind(tag=TAG).error(f"ASR失败: {e}") - return ("", None) - - # 定义声纹识别任务 - def run_voiceprint(): - if not wav_data: - return None - try: - loop = asyncio.new_event_loop() - asyncio.set_event_loop(loop) - try: - # 使用连接的声纹识别提供者 - result = loop.run_until_complete( - conn.voiceprint_provider.identify_speaker(wav_data, conn.session_id) - ) - return result - finally: - loop.close() - except Exception as e: - logger.bind(tag=TAG).error(f"声纹识别失败: {e}") - return None - - # 使用线程池执行器并行运行 - with concurrent.futures.ThreadPoolExecutor(max_workers=2) as thread_executor: - asr_future = thread_executor.submit(run_asr) - - if conn.voiceprint_provider and wav_data: - voiceprint_future = thread_executor.submit(run_voiceprint) - - # 等待两个线程都完成 - asr_result = asr_future.result(timeout=15) - voiceprint_result = voiceprint_future.result(timeout=15) - - results = {"asr": asr_result, "voiceprint": voiceprint_result} - else: - asr_result = asr_future.result(timeout=15) - results = {"asr": asr_result, "voiceprint": None} - - - # 处理结果 - raw_text, _ = results.get("asr", ("", None)) - speaker_name = results.get("voiceprint", None) - - # 记录识别结果 + asr_task = self.speech_to_text(asr_audio_task, conn.session_id, conn.audio_format) + + if conn.voiceprint_provider and wav_data: + voiceprint_task = conn.voiceprint_provider.identify_speaker(wav_data, conn.session_id) + # 并发等待两个结果 + asr_result, voiceprint_result = await asyncio.gather( + asr_task, voiceprint_task, return_exceptions=True + ) + else: + asr_result = await asr_task + voiceprint_result = None + + # 记录识别结果 - 检查是否为异常 + if isinstance(asr_result, Exception): + logger.bind(tag=TAG).error(f"ASR识别失败: {asr_result}") + raw_text = "" + else: + raw_text, _ = asr_result + + if isinstance(voiceprint_result, Exception): + logger.bind(tag=TAG).error(f"声纹识别失败: {voiceprint_result}") + speaker_name = "" + else: + speaker_name = voiceprint_result + if raw_text: logger.bind(tag=TAG).info(f"识别文本: {raw_text}") if speaker_name: logger.bind(tag=TAG).info(f"识别说话人: {speaker_name}") - + # 性能监控 total_time = time.monotonic() - total_start_time logger.bind(tag=TAG).debug(f"总处理耗时: {total_time:.3f}s") - + # 检查文本长度 text_len, _ = remove_punctuation_and_length(raw_text) self.stop_ws_connection() - + if text_len > 0: # 构建包含说话人信息的JSON字符串 enhanced_text = self._build_enhanced_text(raw_text, speaker_name) - + # 使用自定义模块进行上报 await startToChat(conn, enhanced_text) enqueue_asr_report(conn, enhanced_text, asr_audio_task) @@ -241,6 +208,7 @@ class ASRProviderBase(ABC): @staticmethod def decode_opus(opus_data: List[bytes]) -> List[bytes]: """将Opus音频数据解码为PCM数据""" + decoder = None try: decoder = opuslib_next.Decoder(16000, 1) pcm_data = [] @@ -265,3 +233,9 @@ class ASRProviderBase(ABC): except Exception as e: logger.bind(tag=TAG).error(f"音频解码过程发生错误: {e}") return [] + finally: + if decoder is not None: + try: + del decoder + except Exception as e: + logger.bind(tag=TAG).debug(f"释放decoder资源时出错: {e}") diff --git a/main/xiaozhi-server/core/providers/asr/doubao_stream.py b/main/xiaozhi-server/core/providers/asr/doubao_stream.py index 67964075..9c667952 100644 --- a/main/xiaozhi-server/core/providers/asr/doubao_stream.py +++ b/main/xiaozhi-server/core/providers/asr/doubao_stream.py @@ -18,8 +18,6 @@ class ASRProvider(ASRProviderBase): self.interface_type = InterfaceType.STREAM self.config = config self.text = "" - self.max_retries = 3 - self.retry_delay = 2 self.decoder = opuslib_next.Decoder(16000, 1) self.asr_ws = None self.forward_task = None @@ -49,6 +47,8 @@ class ASRProvider(ASRProviderBase): self.channel = config.get("channel", 1) self.auth_method = config.get("auth_method", "token") self.secret = config.get("secret", "access_secret") + end_window_size = config.get("end_window_size") + self.end_window_size = int(end_window_size) if end_window_size else 200 async def open_audio_channels(self, conn): await super().open_audio_channels(conn) @@ -56,14 +56,13 @@ class ASRProvider(ASRProviderBase): async def receive_audio(self, conn, audio, audio_have_voice): conn.asr_audio.append(audio) conn.asr_audio = conn.asr_audio[-10:] - # 存储音频数据 if not hasattr(conn, 'asr_audio_for_voiceprint'): conn.asr_audio_for_voiceprint = [] conn.asr_audio_for_voiceprint.append(audio) - + # 当没有音频数据时处理完整语音片段 - if not audio and len(conn.asr_audio_for_voiceprint) > 0: + if conn.client_listen_mode != "manual" and not audio and len(conn.asr_audio_for_voiceprint) > 0: await self.handle_voice_stop(conn, conn.asr_audio_for_voiceprint) conn.asr_audio_for_voiceprint = [] @@ -179,6 +178,7 @@ class ASRProvider(ASRProviderBase): payload.get("audio_info", {}).get("duration", 0) > 2000 and not utterances and not payload["result"].get("text") + and conn.client_listen_mode != "manual" ): logger.bind(tag=TAG).error(f"识别文本:空") self.text = "" @@ -187,15 +187,44 @@ class ASRProvider(ASRProviderBase): await self.handle_voice_stop(conn, audio_data) break + # 专门处理没有文本的识别结果(手动模式下可能已经识别完成但是没松按键) + elif not payload["result"].get("text") and not utterances: + if conn.client_listen_mode == "manual" and conn.client_voice_stop and len(audio_data) > 0: + logger.bind(tag=TAG).debug("消息结束收到停止信号,触发处理") + await self.handle_voice_stop(conn, audio_data) + # 清理音频缓存 + conn.asr_audio.clear() + conn.reset_vad_states() + break + for utterance in utterances: if utterance.get("definite", False): - self.text = utterance["text"] + current_text = utterance["text"] logger.bind(tag=TAG).info( - f"识别到文本: {self.text}" + f"识别到文本: {current_text}" ) - conn.reset_vad_states() - if len(audio_data) > 15: # 确保有足够音频数据 - await self.handle_voice_stop(conn, audio_data) + + # 手动模式下累积识别结果 + if conn.client_listen_mode == "manual": + if self.text: + self.text += current_text + else: + self.text = current_text + + # 在接收消息中途时收到停止信号 + if conn.client_voice_stop and len(audio_data) > 0: + logger.bind(tag=TAG).debug("消息中途收到停止信号,触发处理") + await self.handle_voice_stop(conn, audio_data) + # 清理音频缓存 + conn.asr_audio.clear() + conn.reset_vad_states() + break + else: + # 自动模式下直接覆盖 + self.text = current_text + conn.reset_vad_states() + if len(audio_data) > 15: # 确保有足够音频数据 + await self.handle_voice_stop(conn, audio_data) break elif "error" in payload: error_msg = payload.get("error", "未知错误") @@ -227,8 +256,6 @@ class ASRProvider(ASRProviderBase): conn.asr_audio_for_voiceprint = [] if hasattr(conn, 'asr_audio'): conn.asr_audio = [] - if hasattr(conn, 'has_valid_voice'): - conn.has_valid_voice = False def stop_ws_connection(self): if self.asr_ws: @@ -236,6 +263,20 @@ class ASRProvider(ASRProviderBase): self.asr_ws = None self.is_processing = False + async def _send_stop_request(self): + """发送最后一个音频帧以通知服务器结束""" + if self.asr_ws: + try: + # 发送结束标记的音频帧(gzip压缩的空数据) + empty_payload = gzip.compress(b"") + last_audio_request = bytearray(self.generate_last_audio_default_header()) + last_audio_request.extend(len(empty_payload).to_bytes(4, "big")) + last_audio_request.extend(empty_payload) + await self.asr_ws.send(last_audio_request) + logger.bind(tag=TAG).debug("已发送结束音频帧") + except Exception as e: + logger.bind(tag=TAG).debug(f"发送结束音频帧时出错: {e}") + def construct_request(self, reqid): req = { "app": { @@ -252,7 +293,7 @@ class ASRProvider(ASRProviderBase): "sequence": 1, "boosting_table_name": self.boosting_table_name, "correct_table_name": self.correct_table_name, - "end_window_size": 200, + "end_window_size": self.end_window_size, }, "audio": { "format": self.format, @@ -370,6 +411,16 @@ class ASRProvider(ASRProviderBase): pass self.forward_task = None self.is_processing = False + + # 显式释放decoder资源 + if hasattr(self, 'decoder') and self.decoder is not None: + try: + del self.decoder + self.decoder = None + logger.bind(tag=TAG).debug("Doubao decoder resources released") + except Exception as e: + logger.bind(tag=TAG).debug(f"释放Doubao decoder资源时出错: {e}") + # 清理所有连接的音频缓冲区 if hasattr(self, '_connections'): for conn in self._connections.values(): @@ -377,5 +428,3 @@ class ASRProvider(ASRProviderBase): conn.asr_audio_for_voiceprint = [] if hasattr(conn, 'asr_audio'): conn.asr_audio = [] - if hasattr(conn, 'has_valid_voice'): - conn.has_valid_voice = False diff --git a/main/xiaozhi-server/core/providers/asr/fun_local.py b/main/xiaozhi-server/core/providers/asr/fun_local.py index 217f17ff..2fd49c36 100644 --- a/main/xiaozhi-server/core/providers/asr/fun_local.py +++ b/main/xiaozhi-server/core/providers/asr/fun_local.py @@ -1,14 +1,16 @@ -import time import os -import sys import io +import sys +import time +import shutil import psutil +import asyncio + from config.logger import setup_logging from typing import Optional, Tuple, List -from core.providers.asr.base import ASRProviderBase from funasr import AutoModel from funasr.utils.postprocess_utils import rich_transcription_postprocess -import shutil +from core.providers.asr.base import ASRProviderBase from core.providers.asr.dto.dto import InterfaceType TAG = __name__ @@ -90,16 +92,17 @@ class ASRProvider(ASRProviderBase): else: file_path = self.save_audio_to_file(pcm_data, session_id) - # 语音识别 + # 语音识别 - 使用线程池避免阻塞事件循环 start_time = time.time() - result = self.model.generate( + result = await asyncio.to_thread( + self.model.generate, input=combined_pcm_data, cache={}, language="auto", use_itn=True, batch_size_s=60, ) - text = rich_transcription_postprocess(result[0]["text"]) + text = await asyncio.to_thread(rich_transcription_postprocess, result[0]["text"]) logger.bind(tag=TAG).debug( f"语音识别耗时: {time.time() - start_time:.3f}s | 结果: {text}" ) diff --git a/main/xiaozhi-server/core/providers/asr/qwen3_asr_flash.py b/main/xiaozhi-server/core/providers/asr/qwen3_asr_flash.py index 84fe1979..51c84a37 100644 --- a/main/xiaozhi-server/core/providers/asr/qwen3_asr_flash.py +++ b/main/xiaozhi-server/core/providers/asr/qwen3_asr_flash.py @@ -1,8 +1,5 @@ import os -import json -import asyncio import tempfile -import difflib from typing import Optional, Tuple, List import dashscope from config.logger import setup_logging @@ -16,7 +13,8 @@ logger = setup_logging() class ASRProvider(ASRProviderBase): def __init__(self, config: dict, delete_audio_file: bool): super().__init__() - self.interface_type = InterfaceType.STREAM + # 音频文件上传类型,流式文本识别输出 + self.interface_type = InterfaceType.NON_STREAM """Qwen3-ASR-Flash ASR初始化""" # 配置参数 @@ -130,27 +128,11 @@ class ASRProvider(ASRProviderBase): # 处理流式响应 full_text = "" - last_text = "" # 用于存储上一个文本片段 for chunk in response: try: text = chunk["output"]["choices"][0]["message"].content[0]["text"] - # 标准化文本片段(去除首尾空格) - normalized_text = text.strip() - # 只有当新文本片段与上一个不同时才处理 - if normalized_text != last_text: - # 提取新增的文本部分 - # 通过比较当前文本和上一个文本,找到新增的部分 - if normalized_text.startswith(last_text): - # 如果当前文本以最后一个文本开头,则新增部分是两者的差集 - new_part = normalized_text[len(last_text):] - else: - # 如果不以最后一个文本开头,说明识别结果发生了较大变化,直接使用当前文本 - new_part = normalized_text - - # 将新增部分添加到完整文本中 - full_text += new_part - last_text = normalized_text - # 这里可以实时处理文本片段,例如通过回调函数 + # 更新为最新的完整文本 + full_text = text.strip() except: pass diff --git a/main/xiaozhi-server/core/providers/asr/xunfei_stream.py b/main/xiaozhi-server/core/providers/asr/xunfei_stream.py index e91c8f09..b7e7886e 100644 --- a/main/xiaozhi-server/core/providers/asr/xunfei_stream.py +++ b/main/xiaozhi-server/core/providers/asr/xunfei_stream.py @@ -5,6 +5,7 @@ import hashlib import asyncio import websockets import opuslib_next +import gc from time import mktime from datetime import datetime from urllib.parse import urlencode @@ -34,9 +35,6 @@ class ASRProvider(ASRProviderBase): self.forward_task = None self.is_processing = False self.server_ready = False - self.last_frame_sent = False # 标记是否已发送最终帧 - self.best_text = "" # 保存最佳识别结果 - self.has_final_result = False # 标记是否收到最终识别结果 # 讯飞配置 self.app_id = config.get("app_id") @@ -51,7 +49,6 @@ class ASRProvider(ASRProviderBase): "domain": config.get("domain", "slm"), "language": config.get("language", "zh_cn"), "accent": config.get("accent", "mandarin"), - "dwa": config.get("dwa", "wpgs"), "result": {"encoding": "utf8", "compress": "raw", "format": "plain"}, } @@ -115,7 +112,7 @@ class ASRProvider(ASRProviderBase): await self._start_recognition(conn) except Exception as e: logger.bind(tag=TAG).error(f"建立ASR连接失败: {str(e)}") - await self._cleanup(conn) + await self._cleanup() return # 发送当前音频数据 @@ -125,7 +122,7 @@ class ASRProvider(ASRProviderBase): await self._send_audio_frame(pcm_frame, STATUS_CONTINUE_FRAME) except Exception as e: logger.bind(tag=TAG).warning(f"发送音频数据时发生错误: {e}") - await self._cleanup(conn) + await self._cleanup() async def _start_recognition(self, conn): """开始识别会话""" @@ -135,6 +132,10 @@ class ASRProvider(ASRProviderBase): ws_url = self.create_url() logger.bind(tag=TAG).info(f"正在连接ASR服务: {ws_url[:50]}...") + # 如果为手动模式,设置超时时长为一分钟 + if conn.client_listen_mode == "manual": + self.iat_params["eos"] = 60000 + self.asr_ws = await websockets.connect( ws_url, max_size=1000000000, @@ -145,8 +146,6 @@ class ASRProvider(ASRProviderBase): logger.bind(tag=TAG).info("ASR WebSocket连接已建立") self.server_ready = False - self.last_frame_sent = False - self.best_text = "" self.forward_task = asyncio.create_task(self._forward_results(conn)) # 发送首帧音频 @@ -195,23 +194,12 @@ class ASRProvider(ASRProviderBase): await self.asr_ws.send(json.dumps(frame_data, ensure_ascii=False)) - # 标记是否发送了最终帧 - if status == STATUS_LAST_FRAME: - self.last_frame_sent = True - logger.bind(tag=TAG).info("标记最终帧已发送") - async def _forward_results(self, conn): """转发识别结果""" try: - while self.asr_ws and not conn.stop_event.is_set(): - # 获取当前连接的音频数据 - audio_data = getattr(conn, "asr_audio_for_voiceprint", []) + while not conn.stop_event.is_set(): try: - # 如果已发送最终帧,增加超时时间等待完整结果 - timeout = 3.0 if self.last_frame_sent else 30.0 - response = await asyncio.wait_for( - self.asr_ws.recv(), timeout=timeout - ) + response = await asyncio.wait_for(self.asr_ws.recv(), timeout=60) result = json.loads(response) logger.bind(tag=TAG).debug(f"收到ASR结果: {result}") @@ -235,144 +223,27 @@ class ASRProvider(ASRProviderBase): # 解码base64文本 decoded_text = base64.b64decode(text_data).decode("utf-8") text_json = json.loads(decoded_text) - # 提取文本内容 text_ws = text_json.get("ws", []) - result_text = "" for i in text_ws: for j in i.get("cw", []): w = j.get("w", "") - result_text += w + self.text += w - # 更新识别文本 - 实时更新策略 - # 只检查是否为空字符串,不再过滤任何标点符号 - # 这样可以确保所有识别到的内容,包括标点符号都能被实时更新 - if result_text and result_text.strip(): - # 实时更新:正常情况下都更新,提高响应速度 - should_update = True - - # 保存最佳文本 - # 1. 如果是识别完成状态或最终帧后收到的结果,优先保存 - # 2. 否则保存最长的有意义文本 - # 取消对标点符号的过滤,只检查是否为空 - # 这样可以保留所有识别到的内容,包括各种标点符号 - is_valid_text = len(result_text.strip()) > 0 - - if ( - self.last_frame_sent or status == 2 - ) and is_valid_text: - self.best_text = result_text - self.has_final_result = True # 标记已收到最终结果 - logger.bind(tag=TAG).debug( - f"保存最终识别结果: {self.best_text}" - ) - elif ( - len(result_text) > len(self.best_text) - and is_valid_text - and not self.has_final_result - ): - self.best_text = result_text - logger.bind(tag=TAG).debug( - f"保存中间最佳文本: {self.best_text}" - ) - - # 如果已发送最终帧,只过滤空文本 - if self.last_frame_sent: - # 只拒绝完全空的结果 - if not result_text.strip(): - should_update = False - logger.bind(tag=TAG).warning( - f"最终帧后拒绝空文本" - ) - - if should_update: - # 处理流式识别结果,避免简单替换导致内容丢失 - # 1. 如果是中间状态(非最终帧后),可能需要替换为更完整的识别 - # 2. 如果是最终帧后收到的结果,可能是对前面文本的补充 - if self.last_frame_sent: - # 最终帧后收到的结果可能是标点符号等补充内容 - # 检查是否需要合并文本而不是替换 - # 如果当前文本是纯标点而前面已有内容,应该追加而不是替换 - if len( - self.text - ) > 0 and result_text.strip() in [ - "。", - ".", - "?", - "?", - "!", - "!", - ",", - ",", - ";", - ";", - ]: - # 对于标点符号,追加到现有文本后 - self.text = ( - self.text.rstrip().rstrip("。.") - + result_text - ) - else: - # 其他情况保持替换逻辑 - self.text = result_text - else: - # 中间状态替换为新的识别结果 - self.text = result_text - - logger.bind(tag=TAG).info( - f"实时更新识别文本: {self.text} (最终帧已发送: {self.last_frame_sent})" - ) - - # 识别完成,但如果还没发送最终帧,继续等待 if status == 2: - logger.bind(tag=TAG).info( - f"识别完成状态已到达,当前识别文本: {self.text}" - ) - - # 如果还没发送最终帧,继续等待 - if not self.last_frame_sent: - logger.bind(tag=TAG).info( - "识别完成但最终帧未发送,继续等待..." - ) - continue - - # 已发送最终帧且收到完成状态,使用最佳策略选择最终结果 - # 优先使用识别完成状态下的最新结果,而不是仅仅基于长度 - if self.best_text: - # 如果当前文本是在最终帧发送后或识别完成状态下收到的,优先使用 - if ( - self.last_frame_sent or status == 2 - ) and self.text.strip(): - logger.bind(tag=TAG).info( - f"使用完成状态下的最新识别结果: {self.text}" - ) - elif len(self.best_text) > len(self.text): - logger.bind(tag=TAG).info( - f"使用更长的最佳文本作为最终结果: {self.text} -> {self.best_text}" - ) - self.text = self.best_text - - logger.bind(tag=TAG).info(f"获取到最终完整文本: {self.text}") + if conn.client_listen_mode == "manual": + audio_data = getattr(conn, 'asr_audio_for_voiceprint', []) + if len(audio_data) > 0: + logger.bind(tag=TAG).debug("收到最终识别结果,触发处理") + await self.handle_voice_stop(conn, audio_data) + # 清理音频缓存 + conn.asr_audio.clear() conn.reset_vad_states() - if len(audio_data) > 15: # 确保有足够音频数据 - # 准备处理结果 - pass break except asyncio.TimeoutError: - if self.last_frame_sent: - # 超时时也使用最佳文本 - if self.best_text and len(self.best_text) > len(self.text): - logger.bind(tag=TAG).info( - f"超时,使用最佳文本: {self.text} -> {self.best_text}" - ) - self.text = self.best_text - logger.bind(tag=TAG).info( - f"最终帧后超时,使用结果: {self.text}" - ) - break - # 如果还没发送最终帧,继续等待 - continue + logger.bind(tag=TAG).error("接收结果超时") + break except websockets.ConnectionClosed: logger.bind(tag=TAG).info("ASR服务连接已关闭") self.is_processing = False @@ -389,17 +260,15 @@ class ASRProvider(ASRProviderBase): if hasattr(e, "__cause__") and e.__cause__: logger.bind(tag=TAG).error(f"错误原因: {str(e.__cause__)}") finally: - if self.asr_ws: - await self.asr_ws.close() - self.asr_ws = None - self.is_processing = False + # 清理连接资源 + await self._cleanup() + + # 清理连接的音频缓存 if conn: if hasattr(conn, "asr_audio_for_voiceprint"): conn.asr_audio_for_voiceprint = [] if hasattr(conn, "asr_audio"): conn.asr_audio = [] - if hasattr(conn, "has_valid_voice"): - conn.has_valid_voice = False async def handle_voice_stop(self, conn, asr_audio_task: List[bytes]): """处理语音停止,发送最后一帧并处理识别结果""" @@ -407,22 +276,13 @@ class ASRProvider(ASRProviderBase): # 先发送最后一帧表示音频结束 if self.asr_ws and self.is_processing: try: - # 取最后一个有效的音频帧作为最后一帧数据 - last_frame = b"" - if asr_audio_task: - last_audio = asr_audio_task[-1] - last_frame = self.decoder.decode(last_audio, 960) - await self._send_audio_frame(last_frame, STATUS_LAST_FRAME) - logger.bind(tag=TAG).info("已发送最后一帧") + await self._send_audio_frame(b"", STATUS_LAST_FRAME) + logger.bind(tag=TAG).debug(f"已发送停止请求") - # 发送最终帧后,给_forward_results适当时间处理最终结果 await asyncio.sleep(0.25) - - logger.bind(tag=TAG).info(f"准备处理最终识别结果: {self.text}") except Exception as e: - logger.bind(tag=TAG).error(f"发送最后一帧失败: {e}") + logger.bind(tag=TAG).error(f"发送停止请求失败: {e}") - # 调用父类的handle_voice_stop方法处理识别结果 await super().handle_voice_stop(conn, asr_audio_task) except Exception as e: logger.bind(tag=TAG).error(f"处理语音停止失败: {e}") @@ -436,40 +296,27 @@ class ASRProvider(ASRProviderBase): self.asr_ws = None self.is_processing = False - async def _cleanup(self, conn): - """清理资源""" - logger.bind(tag=TAG).info( + async def _send_stop_request(self): + """发送停止识别请求(不关闭连接)""" + if self.asr_ws: + try: + # 先停止音频发送 + self.is_processing = False + await self._send_audio_frame(b"", STATUS_LAST_FRAME) + logger.bind(tag=TAG).debug("已发送停止请求") + except Exception as e: + logger.bind(tag=TAG).error(f"发送停止请求失败: {e}") + + async def _cleanup(self): + """清理资源(关闭连接)""" + logger.bind(tag=TAG).debug( f"开始ASR会话清理 | 当前状态: processing={self.is_processing}, server_ready={self.server_ready}" ) - # 发送最后一帧 - if self.asr_ws and self.is_processing: - try: - await self._send_audio_frame(b"", STATUS_LAST_FRAME) - await asyncio.sleep(0.1) - logger.bind(tag=TAG).info("已发送最后一帧") - except Exception as e: - logger.bind(tag=TAG).error(f"发送最后一帧失败: {e}") - # 状态重置 self.is_processing = False self.server_ready = False - self.last_frame_sent = False - self.best_text = "" - self.has_final_result = False - logger.bind(tag=TAG).info("ASR状态已重置") - - # 清理任务 - if self.forward_task and not self.forward_task.done(): - self.forward_task.cancel() - try: - await asyncio.wait_for(self.forward_task, timeout=1.0) - except asyncio.CancelledError: - pass - except Exception as e: - logger.bind(tag=TAG).debug(f"forward_task取消异常: {e}") - finally: - self.forward_task = None + logger.bind(tag=TAG).debug("ASR状态已重置") # 关闭连接 if self.asr_ws: @@ -482,16 +329,10 @@ class ASRProvider(ASRProviderBase): finally: self.asr_ws = None - # 清理连接的音频缓存 - if conn: - if hasattr(conn, "asr_audio_for_voiceprint"): - conn.asr_audio_for_voiceprint = [] - if hasattr(conn, "asr_audio"): - conn.asr_audio = [] - if hasattr(conn, "has_valid_voice"): - conn.has_valid_voice = False + # 清理任务引用 + self.forward_task = None - logger.bind(tag=TAG).info("ASR会话清理完成") + logger.bind(tag=TAG).debug("ASR会话清理完成") async def speech_to_text(self, opus_data, session_id, audio_format): """获取识别结果""" @@ -512,6 +353,16 @@ class ASRProvider(ASRProviderBase): pass self.forward_task = None self.is_processing = False + + # 显式释放decoder资源 + if hasattr(self, 'decoder') and self.decoder is not None: + try: + del self.decoder + self.decoder = None + logger.bind(tag=TAG).debug("Xunfei decoder resources released") + except Exception as e: + logger.bind(tag=TAG).debug(f"释放Xunfei decoder资源时出错: {e}") + # 清理所有连接的音频缓冲区 if hasattr(self, "_connections"): for conn in self._connections.values(): @@ -519,5 +370,3 @@ class ASRProvider(ASRProviderBase): conn.asr_audio_for_voiceprint = [] if hasattr(conn, "asr_audio"): conn.asr_audio = [] - if hasattr(conn, "has_valid_voice"): - conn.has_valid_voice = False diff --git a/main/xiaozhi-server/core/providers/llm/openai/openai.py b/main/xiaozhi-server/core/providers/llm/openai/openai.py index 6e0cd66d..2863c420 100644 --- a/main/xiaozhi-server/core/providers/llm/openai/openai.py +++ b/main/xiaozhi-server/core/providers/llm/openai/openai.py @@ -24,7 +24,6 @@ class LLMProvider(LLMProviderBase): "max_tokens": int, "temperature": lambda x: round(float(x), 1), "top_p": lambda x: round(float(x), 1), - "top_k": int, "frequency_penalty": lambda x: round(float(x), 1), } @@ -40,7 +39,7 @@ class LLMProvider(LLMProviderBase): setattr(self, param, None) logger.debug( - f"意图识别参数初始化: {self.temperature}, {self.max_tokens}, {self.top_p}, {self.top_k}, {self.frequency_penalty}" + f"意图识别参数初始化: {self.temperature}, {self.max_tokens}, {self.top_p}, {self.frequency_penalty}" ) model_key_msg = check_model_key("LLM", self.api_key) @@ -71,7 +70,6 @@ class LLMProvider(LLMProviderBase): "max_tokens": kwargs.get("max_tokens", self.max_tokens), "temperature": kwargs.get("temperature", self.temperature), "top_p": kwargs.get("top_p", self.top_p), - "top_k": kwargs.get("top_k", self.top_k), "frequency_penalty": kwargs.get("frequency_penalty", self.frequency_penalty), } @@ -116,7 +114,6 @@ class LLMProvider(LLMProviderBase): "max_tokens": kwargs.get("max_tokens", self.max_tokens), "temperature": kwargs.get("temperature", self.temperature), "top_p": kwargs.get("top_p", self.top_p), - "top_k": kwargs.get("top_k", self.top_k), "frequency_penalty": kwargs.get("frequency_penalty", self.frequency_penalty), } diff --git a/main/xiaozhi-server/core/providers/memory/base.py b/main/xiaozhi-server/core/providers/memory/base.py index 2ced898d..70c4b409 100644 --- a/main/xiaozhi-server/core/providers/memory/base.py +++ b/main/xiaozhi-server/core/providers/memory/base.py @@ -14,7 +14,7 @@ class MemoryProviderBase(ABC): self.llm = llm @abstractmethod - async def save_memory(self, msgs): + async def save_memory(self, msgs, session_id=None): """Save a new memory for specific role and return memory ID""" print("this is base func", msgs) diff --git a/main/xiaozhi-server/core/providers/memory/mem0ai/mem0ai.py b/main/xiaozhi-server/core/providers/memory/mem0ai/mem0ai.py index 530aa317..7156ab72 100644 --- a/main/xiaozhi-server/core/providers/memory/mem0ai/mem0ai.py +++ b/main/xiaozhi-server/core/providers/memory/mem0ai/mem0ai.py @@ -28,7 +28,7 @@ class MemoryProvider(MemoryProviderBase): logger.bind(tag=TAG).error(f"详细错误: {traceback.format_exc()}") self.use_mem0 = False - async def save_memory(self, msgs): + async def save_memory(self, msgs, session_id=None): if not self.use_mem0: return None if len(msgs) < 2: @@ -41,9 +41,7 @@ class MemoryProvider(MemoryProviderBase): for message in msgs if message.role != "system" ] - result = self.client.add( - messages, user_id=self.role_id - ) + result = self.client.add(messages, user_id=self.role_id) logger.bind(tag=TAG).debug(f"Save memory result: {result}") except Exception as e: logger.bind(tag=TAG).error(f"保存记忆失败: {str(e)}") diff --git a/main/xiaozhi-server/core/providers/memory/mem_local_short/mem_local_short.py b/main/xiaozhi-server/core/providers/memory/mem_local_short/mem_local_short.py index e4486b82..5904d440 100644 --- a/main/xiaozhi-server/core/providers/memory/mem_local_short/mem_local_short.py +++ b/main/xiaozhi-server/core/providers/memory/mem_local_short/mem_local_short.py @@ -4,7 +4,8 @@ import json import os import yaml from config.config_loader import get_project_dir -from config.manage_api_client import save_mem_local_short +from config.manage_api_client import generate_and_save_chat_summary +import asyncio from core.utils.util import check_model_key @@ -74,18 +75,6 @@ short_term_memory_prompt = """ ``` """ -short_term_memory_prompt_only_content = """ -你是一个经验丰富的记忆总结者,擅长将对话内容进行总结摘要,遵循以下规则: -1、总结user的重要信息,以便在未来的对话中提供更个性化的服务 -2、不要重复总结,不要遗忘之前记忆,除非原来的记忆超过了1800字内,否则不要遗忘、不要压缩用户的历史记忆 -3、用户操控的设备音量、播放音乐、天气、退出、不想对话等和用户本身无关的内容,这些信息不需要加入到总结中 -4、聊天内容中的今天的日期时间、今天的天气情况与用户事件无关的数据,这些信息如果当成记忆存储会影响后序对话,这些信息不需要加入到总结中 -5、不要把设备操控的成果结果和失败结果加入到总结中,也不要把用户的一些废话加入到总结中 -6、不要为了总结而总结,如果用户的聊天没有意义,请返回原来的历史记录也是可以的 -7、只需要返回总结摘要,严格控制在1800字内 -8、不要包含代码、xml,不需要解释、注释和说明,保存记忆时仅从对话提取信息,不要混入示例内容 -""" - def extract_json_data(json_code): start = json_code.find("```json") @@ -143,7 +132,7 @@ class MemoryProvider(MemoryProviderBase): with open(self.memory_path, "w", encoding="utf-8") as f: yaml.dump(all_memory, f, allow_unicode=True) - async def save_memory(self, msgs): + async def save_memory(self, msgs, session_id=None): # 打印使用的模型信息 model_info = getattr(self.llm, "model_name", str(self.llm.__class__.__name__)) logger.bind(tag=TAG).debug(f"使用记忆保存模型: {model_info}") @@ -187,14 +176,12 @@ class MemoryProvider(MemoryProviderBase): except Exception as e: print("Error:", e) else: - result = self.llm.response_no_stream( - short_term_memory_prompt_only_content, - msgStr, - max_tokens=2000, - temperature=0.2, - ) - save_mem_local_short(self.role_id, result) - logger.bind(tag=TAG).info(f"Save memory successful - Role: {self.role_id}") + # 当save_to_file为False时,调用Java端的聊天记录总结接口 + summary_id = session_id if session_id else self.role_id + await generate_and_save_chat_summary(summary_id) + logger.bind(tag=TAG).info( + f"Save memory successful - Role: {self.role_id}, Session: {session_id}" + ) return self.short_memory diff --git a/main/xiaozhi-server/core/providers/memory/nomem/nomem.py b/main/xiaozhi-server/core/providers/memory/nomem/nomem.py index 51523be4..6fe0a2be 100644 --- a/main/xiaozhi-server/core/providers/memory/nomem/nomem.py +++ b/main/xiaozhi-server/core/providers/memory/nomem/nomem.py @@ -11,7 +11,7 @@ class MemoryProvider(MemoryProviderBase): def __init__(self, config, summary_memory=None): super().__init__(config) - async def save_memory(self, msgs): + async def save_memory(self, msgs, session_id=None): logger.bind(tag=TAG).debug("nomem mode: No memory saving is performed.") return None diff --git a/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_manager.py b/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_manager.py index 6edc44d1..b1fc0aa3 100644 --- a/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_manager.py +++ b/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_manager.py @@ -3,12 +3,8 @@ import asyncio import os import json -from datetime import timedelta from typing import Dict, Any, List -from mcp import Implementation -from mcp.client.session import SamplingFnT, ElicitationFnT, ListRootsFnT, LoggingFnT, MessageHandlerFnT -from mcp.shared.session import ProgressFnT from mcp.types import LoggingMessageNotificationParams from config.config_loader import get_project_dir @@ -33,6 +29,7 @@ class ServerMCPManager: ) self.clients: Dict[str, ServerMCPClient] = {} self.tools = [] + self._init_lock = asyncio.Lock() def load_config(self) -> Dict[str, Any]: """加载MCP服务配置""" @@ -49,29 +46,50 @@ class ServerMCPManager: ) return {} + async def _init_server(self, name: str, srv_config: Dict[str, Any]): + """初始化单个MCP服务""" + client = None + try: + # 初始化服务端MCP客户端 + logger.bind(tag=TAG).info(f"初始化服务端MCP客户端: {name}") + client = ServerMCPClient(srv_config) + # 设置超时时间10秒 + await asyncio.wait_for(client.initialize(logging_callback=self.logging_callback), timeout=10) + + # 使用锁保护共享状态的修改 + async with self._init_lock: + self.clients[name] = client + client_tools = client.get_available_tools() + self.tools.extend(client_tools) + + except asyncio.TimeoutError: + logger.bind(tag=TAG).error( + f"Failed to initialize MCP server {name}: Timeout" + ) + if client: + await client.cleanup() + except Exception as e: + logger.bind(tag=TAG).error( + f"Failed to initialize MCP server {name}: {e}" + ) + if client: + await client.cleanup() + async def initialize_servers(self) -> None: """初始化所有MCP服务""" config = self.load_config() + tasks = [] for name, srv_config in config.items(): if not srv_config.get("command") and not srv_config.get("url"): logger.bind(tag=TAG).warning( f"Skipping server {name}: neither command nor url specified" ) continue - - try: - # 初始化服务端MCP客户端 - logger.bind(tag=TAG).info(f"初始化服务端MCP客户端: {name}") - client = ServerMCPClient(srv_config) - await client.initialize(logging_callback=self.logging_callback) - self.clients[name] = client - client_tools = client.get_available_tools() - self.tools.extend(client_tools) - - except Exception as e: - logger.bind(tag=TAG).error( - f"Failed to initialize MCP server {name}: {e}" - ) + + tasks.append(self._init_server(name, srv_config)) + + if tasks: + await asyncio.gather(*tasks) # 输出当前支持的服务端MCP工具列表 if hasattr(self.conn, "func_handler") and self.conn.func_handler: diff --git a/main/xiaozhi-server/core/providers/vad/silero.py b/main/xiaozhi-server/core/providers/vad/silero.py index 2263fcb7..81215681 100644 --- a/main/xiaozhi-server/core/providers/vad/silero.py +++ b/main/xiaozhi-server/core/providers/vad/silero.py @@ -36,7 +36,18 @@ class VADProvider(VADProviderBase): # 至少要多少帧才算有语音 self.frame_window_threshold = 3 + def __del__(self): + if hasattr(self, 'decoder') and self.decoder is not None: + try: + del self.decoder + except Exception: + pass + def is_vad(self, conn, opus_packet): + # 手动模式:直接返回True,不进行实时VAD检测,所有音频都缓存 + if conn.client_listen_mode == "manual": + return True + try: pcm_frame = self.decoder.decode(opus_packet, 960) conn.client_audio_buffer.extend(pcm_frame) # 将新数据加入缓冲区 diff --git a/main/xiaozhi-server/core/utils/audioRateController.py b/main/xiaozhi-server/core/utils/audioRateController.py new file mode 100644 index 00000000..01d71b05 --- /dev/null +++ b/main/xiaozhi-server/core/utils/audioRateController.py @@ -0,0 +1,160 @@ +import time +import asyncio +from collections import deque +from config.logger import setup_logging + +TAG = __name__ +logger = setup_logging() + + +class AudioRateController: + """ + 音频速率控制器 - 按照60ms帧时长精确控制音频发送 + 解决高并发下的时间累积误差问题 + """ + + def __init__(self, frame_duration=60): + """ + Args: + frame_duration: 单个音频帧时长(毫秒),默认60ms + """ + self.frame_duration = frame_duration + self.queue = deque() + self.play_position = 0 # 虚拟播放位置(毫秒) + self.start_timestamp = None # 开始时间戳(只读,不修改) + self.pending_send_task = None + self.logger = logger + self.queue_empty_event = asyncio.Event() # 队列清空事件 + self.queue_empty_event.set() # 初始为空状态 + self.queue_has_data_event = asyncio.Event() # 队列数据事件 + + def reset(self): + """重置控制器状态""" + if self.pending_send_task and not self.pending_send_task.done(): + self.pending_send_task.cancel() + # 取消任务后,任务会在下次事件循环时清理,无需阻塞等待 + + self.queue.clear() + self.play_position = 0 + self.start_timestamp = None # 由首个音频包设置 + # 相关事件处理 + self.queue_empty_event.set() + self.queue_has_data_event.clear() + + def add_audio(self, opus_packet): + """添加音频包到队列""" + self.queue.append(("audio", opus_packet)) + # 相关事件处理 + self.queue_empty_event.clear() + self.queue_has_data_event.set() + + def add_message(self, message_callback): + """ + 添加消息到队列(立即发送,不占用播放时间) + + Args: + message_callback: 消息发送回调函数 async def() + """ + self.queue.append(("message", message_callback)) + # 相关事件处理 + self.queue_empty_event.clear() + self.queue_has_data_event.set() + + def _get_elapsed_ms(self): + """获取已经过的时间(毫秒)""" + if self.start_timestamp is None: + return 0 + return (time.monotonic() - self.start_timestamp) * 1000 + + async def check_queue(self, send_audio_callback): + """ + 检查队列并按时发送音频/消息 + + Args: + send_audio_callback: 发送音频的回调函数 async def(opus_packet) + """ + while self.queue: + item = self.queue[0] + item_type = item[0] + + if item_type == "message": + # 消息类型:立即发送,不占用播放时间 + _, message_callback = item + self.queue.popleft() + try: + await message_callback() + except Exception as e: + self.logger.bind(tag=TAG).error(f"发送消息失败: {e}") + raise + + elif item_type == "audio": + if self.start_timestamp is None: + self.start_timestamp = time.monotonic() + + _, opus_packet = item + + # 循环等待直到时间到达 + while True: + # 计算时间差 + elapsed_ms = self._get_elapsed_ms() + output_ms = self.play_position + + if elapsed_ms < output_ms: + # 还不到发送时间,计算等待时长 + wait_ms = output_ms - elapsed_ms + + # 等待后继续检查(允许被中断) + try: + await asyncio.sleep(wait_ms / 1000) + except asyncio.CancelledError: + self.logger.bind(tag=TAG).debug("音频发送任务被取消") + raise + # 等待结束后重新检查时间(循环回到 while True) + else: + # 时间已到,跳出等待循环 + break + + # 时间已到,从队列移除并发送 + self.queue.popleft() + self.play_position += self.frame_duration + try: + await send_audio_callback(opus_packet) + except Exception as e: + self.logger.bind(tag=TAG).error(f"发送音频失败: {e}") + raise + + # 队列处理完后清除事件 + self.queue_empty_event.set() + self.queue_has_data_event.clear() + + def start_sending(self, send_audio_callback): + """ + 启动异步发送任务 + + Args: + send_audio_callback: 发送音频的回调函数 + + Returns: + asyncio.Task: 发送任务 + """ + + async def _send_loop(): + try: + while True: + # 等待队列数据事件,不轮询等待占用CPU + await self.queue_has_data_event.wait() + + await self.check_queue(send_audio_callback) + except asyncio.CancelledError: + self.logger.bind(tag=TAG).debug("音频发送循环已停止") + except Exception as e: + self.logger.bind(tag=TAG).error(f"音频发送循环异常: {e}") + + self.pending_send_task = asyncio.create_task(_send_loop()) + return self.pending_send_task + + def stop_sending(self): + """停止发送任务""" + if self.pending_send_task and not self.pending_send_task.done(): + self.pending_send_task.cancel() + self.logger.bind(tag=TAG).debug("已取消音频发送任务") diff --git a/main/xiaozhi-server/core/utils/cache/config.py b/main/xiaozhi-server/core/utils/cache/config.py index 248c2af7..f85e40bb 100644 --- a/main/xiaozhi-server/core/utils/cache/config.py +++ b/main/xiaozhi-server/core/utils/cache/config.py @@ -19,6 +19,7 @@ class CacheType(Enum): CONFIG = "config" DEVICE_PROMPT = "device_prompt" VOICEPRINT_HEALTH = "voiceprint_health" # 声纹识别健康检查 + AUDIO_DATA = "audio_data" # 音频数据缓存 @dataclass @@ -58,5 +59,8 @@ class CacheConfig: CacheType.VOICEPRINT_HEALTH: cls( strategy=CacheStrategy.TTL, ttl=600, max_size=100 # 10分钟过期 ), + CacheType.AUDIO_DATA: cls( + strategy=CacheStrategy.TTL, ttl=600, max_size=100 # 10分钟过期 + ), } return configs.get(cache_type, cls()) diff --git a/main/xiaozhi-server/core/utils/context_provider.py b/main/xiaozhi-server/core/utils/context_provider.py new file mode 100644 index 00000000..14943d98 --- /dev/null +++ b/main/xiaozhi-server/core/utils/context_provider.py @@ -0,0 +1,64 @@ +import httpx +from typing import Dict, Any, List +from config.logger import setup_logging + +TAG = __name__ + +class ContextDataProvider: + """数据上下文填充,负责从配置的API获取数据""" + + def __init__(self, config: Dict[str, Any], logger=None): + self.config = config + self.logger = logger or setup_logging() + self.context_data = "" + + def fetch_all(self, device_id: str) -> str: + """获取所有配置的上下文数据""" + context_providers = self.config.get("context_providers", []) + if not context_providers: + return "" + + formatted_lines = [] + for provider in context_providers: + url = provider.get("url") + headers = provider.get("headers", {}) + + if not url: + continue + + try: + headers = headers.copy() if isinstance(headers, dict) else {} + # 将 device_id 添加到请求头 + headers["device-id"] = device_id + + # 发送请求 + response = httpx.get(url, headers=headers, timeout=3) + + if response.status_code == 200: + result = response.json() + if isinstance(result, dict): + if result.get("code") == 0: + data = result.get("data") + # 格式化数据 + if isinstance(data, dict): + for k, v in data.items(): + formatted_lines.append(f"- **{k}:** {v}") + elif isinstance(data, list): + for item in data: + formatted_lines.append(f"- {item}") + else: + formatted_lines.append(f"- {data}") + else: + self.logger.bind(tag=TAG).warning(f"API {url} 返回错误码: {result.get('msg')}") + else: + self.logger.bind(tag=TAG).warning(f"API {url} 返回的不是JSON字典") + else: + self.logger.bind(tag=TAG).warning(f"API {url} 请求失败: {response.status_code}") + except Exception as e: + self.logger.bind(tag=TAG).error(f"获取上下文数据 {url} 失败: {e}") + + # 将所有格式化后的行拼接成一个字符串 + self.context_data = "\n".join(formatted_lines) + if self.context_data: + self.logger.bind(tag=TAG).debug(f"已注入动态上下文数据:\n{self.context_data}") + return self.context_data diff --git a/main/xiaozhi-server/core/utils/gc_manager.py b/main/xiaozhi-server/core/utils/gc_manager.py new file mode 100644 index 00000000..e3b958fe --- /dev/null +++ b/main/xiaozhi-server/core/utils/gc_manager.py @@ -0,0 +1,122 @@ +""" +全局GC管理模块 +定期执行垃圾回收,避免频繁触发GC导致的GIL锁问题 +""" + +import gc +import asyncio +import threading +from config.logger import setup_logging + +TAG = __name__ +logger = setup_logging() + + +class GlobalGCManager: + """全局垃圾回收管理器""" + + def __init__(self, interval_seconds=300): + """ + 初始化GC管理器 + + Args: + interval_seconds: GC执行间隔(秒),默认300秒(5分钟) + """ + self.interval_seconds = interval_seconds + self._task = None + self._stop_event = asyncio.Event() + self._lock = threading.Lock() + + async def start(self): + """启动定时GC任务""" + if self._task is not None: + logger.bind(tag=TAG).warning("GC管理器已经在运行") + return + + logger.bind(tag=TAG).info(f"启动全局GC管理器,间隔{self.interval_seconds}秒") + self._stop_event.clear() + self._task = asyncio.create_task(self._gc_loop()) + + async def stop(self): + """停止定时GC任务""" + if self._task is None: + return + + logger.bind(tag=TAG).info("停止全局GC管理器") + self._stop_event.set() + + if self._task and not self._task.done(): + self._task.cancel() + try: + await self._task + except asyncio.CancelledError: + pass + + self._task = None + + async def _gc_loop(self): + """GC循环任务""" + try: + while not self._stop_event.is_set(): + # 等待指定间隔 + try: + await asyncio.wait_for( + self._stop_event.wait(), timeout=self.interval_seconds + ) + # 如果stop_event被设置,退出循环 + break + except asyncio.TimeoutError: + # 超时表示到了执行GC的时间 + pass + + # 执行GC + await self._run_gc() + + except asyncio.CancelledError: + logger.bind(tag=TAG).info("GC循环任务被取消") + raise + except Exception as e: + logger.bind(tag=TAG).error(f"GC循环任务异常: {e}") + finally: + logger.bind(tag=TAG).info("GC循环任务已退出") + + async def _run_gc(self): + """执行垃圾回收""" + try: + # 在线程池中执行GC,避免阻塞事件循环 + loop = asyncio.get_running_loop() + + def do_gc(): + with self._lock: + before = len(gc.get_objects()) + collected = gc.collect() + after = len(gc.get_objects()) + return before, collected, after + + before, collected, after = await loop.run_in_executor(None, do_gc) + logger.bind(tag=TAG).debug( + f"全局GC执行完成 - 回收对象: {collected}, " + f"对象数量: {before} -> {after}" + ) + except Exception as e: + logger.bind(tag=TAG).error(f"执行GC时出错: {e}") + + +# 全局单例 +_gc_manager_instance = None + + +def get_gc_manager(interval_seconds=300): + """ + 获取全局GC管理器实例(单例模式) + + Args: + interval_seconds: GC执行间隔(秒),默认300秒(5分钟) + + Returns: + GlobalGCManager实例 + """ + global _gc_manager_instance + if _gc_manager_instance is None: + _gc_manager_instance = GlobalGCManager(interval_seconds) + return _gc_manager_instance diff --git a/main/xiaozhi-server/core/utils/opus_encoder_utils.py b/main/xiaozhi-server/core/utils/opus_encoder_utils.py index ae7066ce..8d603e22 100644 --- a/main/xiaozhi-server/core/utils/opus_encoder_utils.py +++ b/main/xiaozhi-server/core/utils/opus_encoder_utils.py @@ -102,6 +102,9 @@ class OpusEncoderUtils: def _encode(self, frame: np.ndarray) -> Optional[bytes]: """编码一帧音频数据""" try: + # 编码器已释放,跳过编码 + if not hasattr(self, 'encoder') or self.encoder is None: + return None # 将numpy数组转换为bytes frame_bytes = frame.tobytes() # opuslib要求输入字节数必须是channels*2的倍数 @@ -128,5 +131,9 @@ class OpusEncoderUtils: def close(self): """关闭编码器并释放资源""" - # opuslib没有明确的关闭方法,Python的垃圾回收会处理 - pass \ No newline at end of file + if hasattr(self, 'encoder') and self.encoder: + try: + del self.encoder + self.encoder = None + except Exception as e: + logging.error(f"Error releasing Opus encoder: {e}") \ No newline at end of file diff --git a/main/xiaozhi-server/core/utils/prompt_manager.py b/main/xiaozhi-server/core/utils/prompt_manager.py index 444b16ee..70f14d92 100644 --- a/main/xiaozhi-server/core/utils/prompt_manager.py +++ b/main/xiaozhi-server/core/utils/prompt_manager.py @@ -4,7 +4,6 @@ """ import os -import cnlunar from typing import Dict, Any from config.logger import setup_logging from jinja2 import Template @@ -60,6 +59,11 @@ class PromptManager: self.cache_manager = cache_manager self.CacheType = CacheType + + # 初始化上下文源 + from core.utils.context_provider import ContextDataProvider + self.context_provider = ContextDataProvider(config, self.logger) + self.context_data = {} self._load_base_template() @@ -180,10 +184,33 @@ class PromptManager: def update_context_info(self, conn, client_ip: str): """同步更新上下文信息""" try: - # 获取位置信息(使用全局缓存) - local_address = self._get_location_info(client_ip) - # 获取天气信息(使用全局缓存) - self._get_weather_info(conn, local_address) + local_address = "" + if ( + client_ip + and self.base_prompt_template + and ( + "local_address" in self.base_prompt_template + or "weather_info" in self.base_prompt_template + ) + ): + # 获取位置信息(使用全局缓存) + local_address = self._get_location_info(client_ip) + + if ( + self.base_prompt_template + and "weather_info" in self.base_prompt_template + and local_address + ): + # 获取天气信息(使用全局缓存) + self._get_weather_info(conn, local_address) + + # 获取配置的上下文数据 + if hasattr(conn, "device_id") and conn.device_id: + if self.base_prompt_template and "dynamic_context" in self.base_prompt_template: + self.context_data = self.context_provider.fetch_all(conn.device_id) + else: + self.context_data = "" + self.logger.bind(tag=TAG).debug(f"上下文信息更新完成") except Exception as e: @@ -230,6 +257,7 @@ class PromptManager: emojiList=EMOJI_List, device_id=device_id, client_ip=client_ip, + dynamic_context=self.context_data, *args, **kwargs, ) diff --git a/main/xiaozhi-server/core/utils/util.py b/main/xiaozhi-server/core/utils/util.py index f10cc369..4547ad8f 100644 --- a/main/xiaozhi-server/core/utils/util.py +++ b/main/xiaozhi-server/core/utils/util.py @@ -4,6 +4,7 @@ import json import copy import wave import socket +import asyncio import requests import subprocess import numpy as np @@ -268,56 +269,82 @@ def audio_to_data_stream( pcm_to_data_stream(raw_data, is_opus, callback) -def audio_to_data(audio_file_path: str, is_opus: bool = True) -> list[bytes]: +async def audio_to_data( + audio_file_path: str, is_opus: bool = True, use_cache: bool = True +) -> list[bytes]: """ 将音频文件转换为Opus/PCM编码的帧列表 Args: audio_file_path: 音频文件路径 is_opus: 是否进行Opus编码 + use_cache: 是否使用缓存 """ - # 获取文件后缀名 - file_type = os.path.splitext(audio_file_path)[1] - if file_type: - file_type = file_type.lstrip(".") - # 读取音频文件,-nostdin 参数:不要从标准输入读取数据,否则FFmpeg会阻塞 - audio = AudioSegment.from_file( - audio_file_path, format=file_type, parameters=["-nostdin"] - ) + from core.utils.cache.manager import cache_manager + from core.utils.cache.config import CacheType - # 转换为单声道/16kHz采样率/16位小端编码(确保与编码器匹配) - audio = audio.set_channels(1).set_frame_rate(16000).set_sample_width(2) + # 生成缓存键,包含文件路径和编码类型 + cache_key = f"{audio_file_path}:{is_opus}" - # 获取原始PCM数据(16位小端) - raw_data = audio.raw_data + # 尝试从缓存获取结果 + if use_cache: + cached_result = cache_manager.get(CacheType.AUDIO_DATA, cache_key) + if cached_result is not None: + return cached_result - # 初始化Opus编码器 - encoder = opuslib_next.Encoder(16000, 1, opuslib_next.APPLICATION_AUDIO) + def _sync_audio_to_data(): + # 获取文件后缀名 + file_type = os.path.splitext(audio_file_path)[1] + if file_type: + file_type = file_type.lstrip(".") + # 读取音频文件,-nostdin 参数:不要从标准输入读取数据,否则FFmpeg会阻塞 + audio = AudioSegment.from_file( + audio_file_path, format=file_type, parameters=["-nostdin"] + ) - # 编码参数 - frame_duration = 60 # 60ms per frame - frame_size = int(16000 * frame_duration / 1000) # 960 samples/frame + # 转换为单声道/16kHz采样率/16位小端编码(确保与编码器匹配) + audio = audio.set_channels(1).set_frame_rate(16000).set_sample_width(2) - datas = [] - # 按帧处理所有音频数据(包括最后一帧可能补零) - for i in range(0, len(raw_data), frame_size * 2): # 16bit=2bytes/sample - # 获取当前帧的二进制数据 - chunk = raw_data[i : i + frame_size * 2] + # 获取原始PCM数据(16位小端) + raw_data = audio.raw_data - # 如果最后一帧不足,补零 - if len(chunk) < frame_size * 2: - chunk += b"\x00" * (frame_size * 2 - len(chunk)) + # 初始化Opus编码器 + encoder = opuslib_next.Encoder(16000, 1, opuslib_next.APPLICATION_AUDIO) - if is_opus: - # 转换为numpy数组处理 - np_frame = np.frombuffer(chunk, dtype=np.int16) - # 编码Opus数据 - frame_data = encoder.encode(np_frame.tobytes(), frame_size) - else: - frame_data = chunk if isinstance(chunk, bytes) else bytes(chunk) + # 编码参数 + frame_duration = 60 # 60ms per frame + frame_size = int(16000 * frame_duration / 1000) # 960 samples/frame - datas.append(frame_data) + datas = [] + # 按帧处理所有音频数据(包括最后一帧可能补零) + for i in range(0, len(raw_data), frame_size * 2): # 16bit=2bytes/sample + # 获取当前帧的二进制数据 + chunk = raw_data[i : i + frame_size * 2] - return datas + # 如果最后一帧不足,补零 + if len(chunk) < frame_size * 2: + chunk += b"\x00" * (frame_size * 2 - len(chunk)) + + if is_opus: + # 转换为numpy数组处理 + np_frame = np.frombuffer(chunk, dtype=np.int16) + # 编码Opus数据 + frame_data = encoder.encode(np_frame.tobytes(), frame_size) + else: + frame_data = chunk if isinstance(chunk, bytes) else bytes(chunk) + + datas.append(frame_data) + + return datas + + loop = asyncio.get_running_loop() + # 在单独的线程中执行同步的音频处理操作 + result = await loop.run_in_executor(None, _sync_audio_to_data) + + # 将结果存入缓存,使用配置中定义的TTL(10分钟) + if use_cache: + cache_manager.set(CacheType.AUDIO_DATA, cache_key, result) + + return result def audio_bytes_to_data_stream( @@ -372,26 +399,33 @@ def opus_datas_to_wav_bytes(opus_datas, sample_rate=16000, channels=1): 将opus帧列表解码为wav字节流 """ decoder = opuslib_next.Decoder(sample_rate, channels) - pcm_datas = [] + try: + pcm_datas = [] - frame_duration = 60 # ms - frame_size = int(sample_rate * frame_duration / 1000) # 960 + frame_duration = 60 # ms + frame_size = int(sample_rate * frame_duration / 1000) # 960 - for opus_frame in opus_datas: - # 解码为PCM(返回bytes,2字节/采样点) - pcm = decoder.decode(opus_frame, frame_size) - pcm_datas.append(pcm) + for opus_frame in opus_datas: + # 解码为PCM(返回bytes,2字节/采样点) + pcm = decoder.decode(opus_frame, frame_size) + pcm_datas.append(pcm) - pcm_bytes = b"".join(pcm_datas) + pcm_bytes = b"".join(pcm_datas) - # 写入wav字节流 - wav_buffer = BytesIO() - with wave.open(wav_buffer, "wb") as wf: - wf.setnchannels(channels) - wf.setsampwidth(2) # 16bit - wf.setframerate(sample_rate) - wf.writeframes(pcm_bytes) - return wav_buffer.getvalue() + # 写入wav字节流 + wav_buffer = BytesIO() + with wave.open(wav_buffer, "wb") as wf: + wf.setnchannels(channels) + wf.setsampwidth(2) # 16bit + wf.setframerate(sample_rate) + wf.writeframes(pcm_bytes) + return wav_buffer.getvalue() + finally: + if decoder is not None: + try: + del decoder + except Exception: + pass def check_vad_update(before_config, new_config): diff --git a/main/xiaozhi-server/core/websocket_server.py b/main/xiaozhi-server/core/websocket_server.py index 46b60a71..168b436e 100644 --- a/main/xiaozhi-server/core/websocket_server.py +++ b/main/xiaozhi-server/core/websocket_server.py @@ -1,10 +1,37 @@ import asyncio -import json +import logging import websockets from config.logger import setup_logging + + +class SuppressInvalidHandshakeFilter(logging.Filter): + """过滤掉无效握手错误日志(如HTTPS访问WS端口)""" + + def filter(self, record): + msg = record.getMessage() + suppress_keywords = [ + "opening handshake failed", + "did not receive a valid HTTP request", + "connection closed while reading HTTP request", + "line without CRLF", + ] + return not any(keyword in msg for keyword in suppress_keywords) + + +def _setup_websockets_logger(): + """配置 websockets 相关的所有 logger,过滤无效握手错误""" + filter_instance = SuppressInvalidHandshakeFilter() + for logger_name in ["websockets", "websockets.server", "websockets.client"]: + logger = logging.getLogger(logger_name) + logger.addFilter(filter_instance) + + +_setup_websockets_logger() + + from core.connection import ConnectionHandler -from config.config_loader import get_config_from_api +from config.config_loader import get_config_from_api_async from core.auth import AuthManager, AuthenticationError from core.utils.modules_initialize import initialize_modules from core.utils.util import check_vad_update, check_asr_update @@ -133,8 +160,8 @@ class WebSocketServer: """ try: async with self.config_lock: - # 重新获取配置 - new_config = get_config_from_api(self.config) + # 重新获取配置(使用异步版本) + new_config = await get_config_from_api_async(self.config) if new_config is None: self.logger.bind(tag=TAG).error("获取新配置失败") return False diff --git a/main/xiaozhi-server/performance_tester/performance_tester_stream_asr.py b/main/xiaozhi-server/performance_tester/performance_tester_stream_asr.py index f4508ef2..f81f6317 100644 --- a/main/xiaozhi-server/performance_tester/performance_tester_stream_asr.py +++ b/main/xiaozhi-server/performance_tester/performance_tester_stream_asr.py @@ -50,7 +50,9 @@ class BaseASRTester: raise NotImplementedError def _calculate_result(self, service_name, latencies, test_count): - valid_latencies = [l for l in latencies if l > 0] + """计算测试结果(修复:正确处理None值,剔除失败测试)""" + # 剔除None值(失败的测试)和无效延迟,只统计有效延迟 + valid_latencies = [l for l in latencies if l is not None and l > 0] if valid_latencies: avg_latency = sum(valid_latencies) / len(valid_latencies) status = f"成功({len(valid_latencies)}/{test_count}次有效)" @@ -64,16 +66,45 @@ class DoubaoStreamASRTester(BaseASRTester): def __init__(self): super().__init__("DoubaoStreamASR") - def _generate_header(self): + def _generate_header( + self, + version=0x01, + message_type=0x01, + message_type_specific_flags=0x00, + serial_method=0x01, + compression_type=0x01, + reserved_data=0x00, + extension_header: bytes = b"", + ): + """生成协议头(修复:使用正确的Header格式)""" header = bytearray() - header.append((0x01 << 4) | 0x01) - header.append((0x01 << 4) | 0x00) - header.append((0x01 << 4) | 0x01) - header.append(0x00) + header_size = int(len(extension_header) / 4) + 1 + header.append((version << 4) | header_size) + header.append((message_type << 4) | message_type_specific_flags) + header.append((serial_method << 4) | compression_type) + header.append(reserved_data) + header.extend(extension_header) return header def _generate_audio_default_header(self): - return self._generate_header() + """生成音频数据Header""" + return self._generate_header( + version=0x01, + message_type=0x02, + message_type_specific_flags=0x00, # 普通音频帧 + serial_method=0x01, + compression_type=0x01, + ) + + def _generate_last_audio_header(self): + """生成最后一帧音频的Header(标记音频结束)""" + return self._generate_header( + version=0x01, + message_type=0x02, + message_type_specific_flags=0x02, # 0x02表示这是最后一帧 + serial_method=0x01, + compression_type=0x01, + ) def _parse_response(self, res: bytes) -> dict: try: @@ -110,6 +141,7 @@ class DoubaoStreamASRTester(BaseASRTester): ws_url = "wss://openspeech.bytedance.com/api/v3/sauc/bigmodel" appid = self.asr_config["appid"] access_token = self.asr_config["access_token"] + cluster = self.asr_config.get("cluster", "volcengine_input_common") uid = self.asr_config.get("uid", "streaming_asr_service") start_time = time.time() @@ -130,7 +162,7 @@ class DoubaoStreamASRTester(BaseASRTester): close_timeout=10 ) as ws: request_params = { - "app": {"appid": appid, "token": access_token}, + "app": {"appid": appid, "cluster": cluster, "token": access_token}, "user": {"uid": uid}, "request": { "reqid": str(uuid.uuid4()), @@ -166,8 +198,9 @@ class DoubaoStreamASRTester(BaseASRTester): if audio_data.startswith(b'RIFF'): audio_data = audio_data[44:] + # 发送音频数据(使用最后一帧标记,告诉服务端音频已结束) payload = gzip.compress(audio_data) - audio_request = bytearray(self._generate_audio_default_header()) + audio_request = bytearray(self._generate_last_audio_header()) # 修复:使用最后一帧Header audio_request.extend(len(payload).to_bytes(4, "big")) audio_request.extend(payload) await ws.send(audio_request) @@ -175,11 +208,12 @@ class DoubaoStreamASRTester(BaseASRTester): first_chunk = await ws.recv() latency = time.time() - start_time latencies.append(latency) + print(f"[豆包ASR] 第{i+1}次 首词延迟: {latency:.3f}s") await ws.close() except Exception as e: print(f"[豆包ASR] 第{i+1}次测试失败: {str(e)}") - latencies.append(0) + latencies.append(None) return self._calculate_result("豆包流式ASR", latencies, test_count) @@ -189,11 +223,12 @@ class QwenASRFlashTester(BaseASRTester): super().__init__("Qwen3ASRFlash") async def _test_single(self, audio_file_info): - start_time = time.time() temp_file_path = None try: audio_data = audio_file_info['data'] + + # 优化:将临时文件准备工作移到计时前,减少磁盘IO对性能测试的影响 with tempfile.NamedTemporaryFile(suffix='.wav', delete=False) as f: temp_file_path = f.name @@ -221,6 +256,9 @@ class QwenASRFlashTester(BaseASRTester): dashscope.api_key = api_key + # 统一计时起点:在API调用前开始计时(但文件准备已完成) + start_time = time.time() + response = dashscope.MultiModalConversation.call( model="qwen3-asr-flash", messages=messages, @@ -257,10 +295,10 @@ class QwenASRFlashTester(BaseASRTester): # print(f"\n[通义ASR] 开始第 {i+1} 次测试...") latency = await self._test_single(self.test_audio_files[0]) latencies.append(latency) - # print(f"[通义ASR] 第{i+1}次成功 延迟: {latency:.3f}s") + print(f"[通义ASR] 第{i+1}次 首词延迟: {latency:.3f}s") except Exception as e: # print(f"[通义ASR] 第{i+1}次测试失败: {str(e)}") - latencies.append(0) + latencies.append(None) return self._calculate_result("通义千问ASR", latencies, test_count) @@ -268,134 +306,115 @@ class QwenASRFlashTester(BaseASRTester): class XunfeiStreamASRTester(BaseASRTester): def __init__(self): super().__init__("XunfeiStreamASR") - + def _create_url(self): - """生成讯飞ASR认证URL""" - url = 'ws://iat.cn-huabei-1.xf-yun.com/v1' - # 生成RFC1123格式的时间戳 + url = "wss://iat-api.xfyun.cn/v2/iat" now = datetime.now() date = format_date_time(mktime(now.timetuple())) - # 拼接字符串 - signature_origin = "host: " + "iat.cn-huabei-1.xf-yun.com" + "\n" - signature_origin += "date: " + date + "\n" - signature_origin += "GET " + "/v1 " + "HTTP/1.1" + signature_origin = f"host: iat-api.xfyun.cn\ndate: {date}\nGET /v2/iat HTTP/1.1" + signature_sha = hmac.new( + self.asr_config["api_secret"].encode('utf-8'), + signature_origin.encode('utf-8'), + hashlib.sha256 + ).digest() + signature_sha = base64.b64encode(signature_sha).decode() - # 进行hmac-sha256进行加密 - signature_sha = hmac.new(self.asr_config["api_secret"].encode('utf-8'), signature_origin.encode('utf-8'), - digestmod=hashlib.sha256).digest() - signature_sha = base64.b64encode(signature_sha).decode(encoding='utf-8') + authorization_origin = f'api_key="{self.asr_config["api_key"]}", algorithm="hmac-sha256", headers="host date request-line", signature="{signature_sha}"' + authorization = base64.b64encode(authorization_origin.encode()).decode() - authorization_origin = "api_key=\"%s\", algorithm=\"%s\", headers=\"%s\", signature=\"%s\"" % ( - self.asr_config["api_key"], "hmac-sha256", "host date request-line", signature_sha) - authorization = base64.b64encode(authorization_origin.encode('utf-8')).decode(encoding='utf-8') + v = {"authorization": authorization, "date": date, "host": "iat-api.xfyun.cn"} + return url + "?" + parse.urlencode(v) - # 将请求的鉴权参数组合为字典 - v = { - "authorization": authorization, - "date": date, - "host": "iat.cn-huabei-1.xf-yun.com" - } - - # 拼接鉴权参数,生成url - url = url + '?' + parse.urlencode(v) - return url - - async def test(self, test_count=5): + async def test(self, test_count: int = 5): if not self.test_audio_files: return {"name": "讯飞流式ASR", "latency": 0, "status": "失败: 未找到测试音频"} if not self.asr_config: return {"name": "讯飞流式ASR", "latency": 0, "status": "失败: 未配置"} - - # 检查必要的配置参数 - required_keys = ["app_id", "api_key", "api_secret"] - for key in required_keys: - if key not in self.asr_config: - return {"name": "讯飞流式ASR", "latency": 0, "status": f"失败: 缺少配置项 {key}"} - + + required = ["app_id", "api_key", "api_secret"] + for k in required: + if k not in self.asr_config: + return {"name": "讯飞流式ASR", "latency": 0, "status": f"失败: 缺少配置 {k}"} + latencies = [] - STATUS_FIRST_FRAME = 0 - + frame_size = 1280 + audio_raw = self.test_audio_files[0]['data'] + if audio_raw.startswith(b'RIFF'): + audio_raw = audio_raw[44:] + for i in range(test_count): try: - # 生成认证URL - ws_url = self._create_url() - - # 获取音频数据 - audio_data = self.test_audio_files[0]['data'] - if audio_data.startswith(b'RIFF'): - audio_data = audio_data[44:] # 跳过WAV文件头 - - # 识别参数 - iat_params = { - "domain": self.asr_config.get("domain", "slm"), - "language": self.asr_config.get("language", "zh_cn"), - "accent": self.asr_config.get("accent", "mandarin"), - "dwa": self.asr_config.get("dwa", "wpgs"), - "result": { - "encoding": "utf8", - "compress": "raw", - "format": "plain" - } - } - - # 准备首帧数据 - first_frame_data = { - "header": { - "status": STATUS_FIRST_FRAME, - "app_id": self.asr_config["app_id"] - }, - "parameter": { - "iat": iat_params - }, - "payload": { - "audio": { - "audio": base64.b64encode(audio_data[:960]).decode('utf-8'), - "sample_rate": 16000, - "encoding": "raw" - } - } - } - - # 启动连接并测量时间 start_time = time.time() - + ws_url = self._create_url() + async with websockets.connect( ws_url, - max_size=1000000000, + additional_headers={"User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64)"}, + max_size=1 << 30, ping_interval=None, ping_timeout=None, close_timeout=30, ) as ws: - # 发送首帧数据 - await ws.send(json.dumps(first_frame_data, ensure_ascii=False)) - print(f"[讯飞ASR] 第{i+1}次测试:已发送首帧,等待响应...") - - # 直接等待第一个响应并计算延迟 - # 参考豆包和通义千问的实现方式,简化逻辑 - response_received = False - while not response_received: - try: - # 设置较大的超时时间 - response = await asyncio.wait_for(ws.recv(), timeout=30.0) - - # 收到响应立即计算延迟,不管内容是什么 - # 这样可以准确测量首包到达时间 - latency = time.time() - start_time - latencies.append(latency) - response_received = True - - print(f"[讯飞ASR] 第{i+1}次测试:收到首包响应,延迟: {latency:.3f}s") + + # 第一帧:移除 punc 字段,避免未知参数错误 + await ws.send(json.dumps({ + "common": {"app_id": self.asr_config["app_id"]}, + "business": { + "domain": "iat", + "language": "zh_cn", + "accent": "mandarin", + "dwa": "wpgs", + "vad_eos": 5000 + # 已移除 "punc": True + }, + "data": { + "status": 0, + "format": "audio/L16;rate=16000", + "encoding": "raw", + "audio": base64.b64encode(audio_raw[:frame_size]).decode() + } + }, ensure_ascii=False)) + + # 后续所有帧 + pos = frame_size + while pos < len(audio_raw): + chunk = audio_raw[pos:pos + frame_size] + status = 2 if (pos + frame_size >= len(audio_raw)) else 1 + await ws.send(json.dumps({ + "data": { + "status": status, + "format": "audio/L16;rate=16000", + "encoding": "raw", + "audio": base64.b64encode(chunk).decode() + } + }, ensure_ascii=False)) + if status == 2: break - except asyncio.TimeoutError: - print(f"[讯飞ASR] 第{i+1}次测试:响应超时") - raise Exception("获取响应超时") + pos += frame_size + + # 接收首词 + first_token = True + async for message in ws: + data = json.loads(message) + if data.get("code") != 0: + raise Exception(f"讯飞错误: {data.get('message')}") + + ws_result = data.get("data", {}).get("result", {}).get("ws") + if ws_result: + text = "".join(cw.get("w", "") for seg in ws_result for cw in seg.get("cw", [])) + if text.strip() and first_token: + latency = time.time() - start_time + latencies.append(latency) + print(f"[讯飞ASR] 第{i+1}次 首词延迟: {latency:.3f}s") + first_token = False + break + except Exception as e: print(f"[讯飞ASR] 第{i+1}次测试失败: {str(e)}") - latencies.append(0) - - return self._calculate_result("讯飞流式ASR", latencies, test_count) + latencies.append(None) + return self._calculate_result("讯飞流式ASR", latencies, test_count) class ASRPerformanceSuite: def __init__(self): self.testers = [] @@ -438,8 +457,9 @@ class ASRPerformanceSuite: print(tabulate(table_data, headers=["ASR服务", "首词延迟", "状态"], tablefmt="grid")) print("\n测试说明:") - print("- 测量从发送请求到接收第一个有效识别文本的时间") - print("- 超时控制: DashScope 默认超时,豆包 WebSocket 超时10秒") + print("- 计时起点: 建立连接前(包含握手、发送音频、接收首个识别结果全流程)") + print("- 通义千问优化: 临时文件准备在计时前完成,减少磁盘IO对测试的影响") + print("- 错误处理: 失败的测试不计入平均值,只统计成功测试的延迟") print("- 排序规则: 成功的按延迟升序,失败的排在后面") async def run(self, test_count=5): diff --git a/main/xiaozhi-server/performance_tester/performance_tester_stream_tts.py b/main/xiaozhi-server/performance_tester/performance_tester_stream_tts.py index ce53fea5..31d49a3e 100644 --- a/main/xiaozhi-server/performance_tester/performance_tester_stream_tts.py +++ b/main/xiaozhi-server/performance_tester/performance_tester_stream_tts.py @@ -35,11 +35,12 @@ class StreamTTSPerformanceTester: host = tts_config["host"] ws_url = f"wss://{host}/ws/v1" + # 统一计时起点:在建立连接前开始计时 start_time = time.time() async with websockets.connect(ws_url, extra_headers={"X-NLS-Token": token}) as ws: task_id = str(uuid.uuid4()) message_id = str(uuid.uuid4()) - + start_request = { "header": { "message_id": message_id, @@ -55,14 +56,15 @@ class StreamTTSPerformanceTester: "volume": 50, "speech_rate": 0, "pitch_rate": 0, + "enable_subtitle": True, } } await ws.send(json.dumps(start_request)) - + start_response = json.loads(await ws.recv()) if start_response["header"]["name"] != "SynthesisStarted": raise Exception("启动合成失败") - + run_request = { "header": { "message_id": str(uuid.uuid4()), @@ -74,23 +76,142 @@ class StreamTTSPerformanceTester: "payload": {"text": text} } await ws.send(json.dumps(run_request)) - + while True: response = await ws.recv() if isinstance(response, bytes): latency = time.time() - start_time latencies.append(latency) + print(f"[阿里云TTS] 第{i+1}次 首词延迟: {latency:.3f}s") break elif isinstance(response, str): data = json.loads(response) if data["header"]["name"] == "TaskFailed": raise Exception(f"合成失败: {data['payload']['error_info']}") - + except Exception as e: - latencies.append(0) + print(f"[阿里云TTS] 第{i+1}次测试失败: {str(e)}") + latencies.append(None) return self._calculate_result("阿里云TTS", latencies, test_count) + async def test_alibl_tts(self, text=None, test_count=5): + """测试阿里云百炼CosyVoice流式TTS首词延迟""" + text = text or self.test_texts[0] + latencies = [] + + for i in range(test_count): + try: + tts_config = self.config["TTS"]["AliBLTTS"] + api_key = tts_config["api_key"] + model = tts_config.get("model", "cosyvoice-v2") + voice = tts_config.get("voice", "longxiaochun_v2") + format_type = tts_config.get("format", "pcm") + sample_rate = int(tts_config.get("sample_rate", "24000")) + + ws_url = "wss://dashscope.aliyuncs.com/api-ws/v1/inference/" + headers = { + "Authorization": f"Bearer {api_key}", + "X-DashScope-DataInspection": "enable", + } + + start_time = time.time() + + async with websockets.connect( + ws_url, + additional_headers=headers, + ping_interval=30, + ping_timeout=10, + close_timeout=10, + max_size=10 * 1024 * 1024, + ) as ws: + session_id = uuid.uuid4().hex + + # 1. 发送 run-task(启动任务) + run_task_message = { + "header": { + "action": "run-task", + "task_id": session_id, + "streaming": "duplex", + }, + "payload": { + "task_group": "audio", + "task": "tts", + "function": "SpeechSynthesizer", + "model": model, + "parameters": { + "text_type": "PlainText", + "voice": voice, + "format": format_type, + "sample_rate": sample_rate, + "volume": 50, + "rate": 1.0, + "pitch": 1.0, + }, + "input": {} + }, + } + await ws.send(json.dumps(run_task_message)) + + # 2. 等待 task-started 事件(关键!必须等这个再发文本) + task_started = False + while not task_started: + msg = await ws.recv() + if isinstance(msg, str): + data = json.loads(msg) + header = data.get("header", {}) + event = header.get("event") + if event == "task-started": + task_started = True + print(f"[阿里云百炼TTS] 第{i+1}次 任务启动成功") + elif event == "task-failed": + raise Exception(f"启动失败: {header.get('error_message', '未知错误')}") + + # 3. 发送 continue-task(发送文本!这是正确动作) + continue_task_message = { + "header": { + "action": "continue-task", # 改回 continue-task + "task_id": session_id, + "streaming": "duplex", + }, + "payload": {"input": {"text": text}}, + } + await ws.send(json.dumps(continue_task_message)) + + # 4. 发送 finish-task(结束任务) + finish_task_message = { + "header": { + "action": "finish-task", + "task_id": session_id, + "streaming": "duplex", + }, + "payload": {"input": {}} + } + await ws.send(json.dumps(finish_task_message)) + + # 5. 等待第一个音频数据块 + while True: + msg = await asyncio.wait_for(ws.recv(), timeout=15.0) + if isinstance(msg, (bytes, bytearray)) and len(msg) > 0: + latency = time.time() - start_time + print(f"[阿里云百炼TTS] 第{i+1}次 首词延迟: {latency:.3f}s") + latencies.append(latency) + break + elif isinstance(msg, str): + data = json.loads(msg) + event = data.get("header", {}).get("event") + if event == "task-failed": + raise Exception(f"合成失败: {data}") + elif event == "task-finished": + if not latencies or latencies[-1] is None: + raise Exception("任务结束但未收到音频") + + except Exception as e: + print(f"[阿里云百炼TTS] 第{i+1}次失败: {str(e)}") + latencies.append(None) + + return self._calculate_result("阿里云百炼TTS", latencies, test_count) + async def test_doubao_tts(self, text=None, test_count=5): """测试火山引擎流式TTS首词延迟(测试多次取平均)""" text = text or self.test_texts[0] @@ -114,13 +235,12 @@ class StreamTTSPerformanceTester: } async with websockets.connect(ws_url, additional_headers=ws_header, max_size=1000000000) as ws: session_id = uuid.uuid4().hex - + # 发送会话启动请求 header = bytes([ - (0b0001 << 4) | 0b0001, - 0b0001 << 4 | 0b100, - 0b0001 << 4 | 0b0000, - 0 + (0b0001 << 4) | 0b0001, + 0b0001 << 4 | 0b1011, + 0b0001 << 4 | 0b0000, ]) optional = bytearray() optional.extend((1).to_bytes(4, "big", signed=True)) @@ -129,13 +249,13 @@ class StreamTTSPerformanceTester: optional.extend(session_id_bytes) payload = json.dumps({"speaker": speaker}).encode() await ws.send(header + optional + len(payload).to_bytes(4, "big", signed=True) + payload) - + # 发送文本 header = bytes([ - (0b0001 << 4) | 0b0001, - 0b0001 << 4 | 0b100, - 0b0001 << 4 | 0b0000, - 0 + (0b0001 << 4) | 0b0001, + 0b0001 << 4 | 0b1011, + 0b0001 << 4 | 0b0000, + 0 ]) optional = bytearray() optional.extend((200).to_bytes(4, "big", signed=True)) @@ -144,13 +264,15 @@ class StreamTTSPerformanceTester: optional.extend(session_id_bytes) payload = json.dumps({"text": text, "speaker": speaker}).encode() await ws.send(header + optional + len(payload).to_bytes(4, "big", signed=True) + payload) - + first_chunk = await ws.recv() latency = time.time() - start_time latencies.append(latency) - + print(f"[火山引擎TTS] 第{i+1}次 首词延迟: {latency:.3f}s") + except Exception as e: - latencies.append(0) + print(f"[火山引擎TTS] 第{i+1}次测试失败: {str(e)}") + latencies.append(None) return self._calculate_result("火山引擎TTS", latencies, test_count) @@ -191,22 +313,24 @@ class StreamTTSPerformanceTester: first_chunk = await ws.recv() latency = time.time() - start_time latencies.append(latency) - + print(f"[PaddleSpeechTTS] 第{i+1}次 首词延迟: {latency:.3f}s") + # 发送结束请求 end_request = { "task": "tts", "signal": "end" } await ws.send(json.dumps(end_request)) - + # 确保连接正常关闭 try: await ws.recv() except websockets.exceptions.ConnectionClosedOK: pass - + except Exception as e: - latencies.append(0) + print(f"[PaddleSpeechTTS] 第{i+1}次测试失败: {str(e)}") + latencies.append(None) return self._calculate_result("PaddleSpeechTTS", latencies, test_count) @@ -220,29 +344,32 @@ class StreamTTSPerformanceTester: tts_config = self.config["TTS"]["IndexStreamTTS"] api_url = tts_config.get("api_url") voice = tts_config.get("voice") - + + # 统一计时起点:在建立连接前开始计时 start_time = time.time() - + async with aiohttp.ClientSession() as session: payload = {"text": text, "character": voice} async with session.post(api_url, json=payload, timeout=10) as resp: if resp.status != 200: raise Exception(f"请求失败: {resp.status}, {await resp.text()}") - + async for chunk in resp.content.iter_any(): data = chunk[0] if isinstance(chunk, (list, tuple)) else chunk if not data: continue - + latency = time.time() - start_time latencies.append(latency) + print(f"[IndexStreamTTS] 第{i+1}次 首词延迟: {latency:.3f}s") resp.close() break else: - latencies.append(0) - + latencies.append(None) + except Exception as e: - latencies.append(0) + print(f"[IndexStreamTTS] 第{i+1}次测试失败: {str(e)}") + latencies.append(None) return self._calculate_result("IndexStreamTTS", latencies, test_count) @@ -257,7 +384,8 @@ class StreamTTSPerformanceTester: api_url = tts_config["api_url"] access_token = tts_config["access_token"] voice = tts_config["voice"] - + + # 统一计时起点:在建立连接前开始计时 start_time = time.time() async with aiohttp.ClientSession() as session: params = { @@ -273,21 +401,23 @@ class StreamTTSPerformanceTester: "Authorization": f"Bearer {access_token}", "Content-Type": "application/json", } - + async with session.get(api_url, params=params, headers=headers, timeout=10) as resp: if resp.status != 200: raise Exception(f"请求失败: {resp.status}, {await resp.text()}") - + # 接收第一个数据块 async for _ in resp.content.iter_any(): latency = time.time() - start_time latencies.append(latency) + print(f"[LinkeraiTTS] 第{i+1}次 首词延迟: {latency:.3f}s") break else: - latencies.append(0) - + latencies.append(None) + except Exception as e: - latencies.append(0) + print(f"[LinkeraiTTS] 第{i+1}次测试失败: {str(e)}") + latencies.append(None) return self._calculate_result("LinkeraiTTS", latencies, test_count) @@ -305,10 +435,9 @@ class StreamTTSPerformanceTester: api_secret = tts_config["api_secret"] api_url = tts_config.get("api_url", "wss://cbm01.cn-huabei-1.xf-yun.com/v1/private/mcd9m97e6") voice = tts_config.get("voice", "x5_lingxiaoxuan_flow") - # 生成认证URL auth_url = self._create_xunfei_auth_url(api_key, api_secret, api_url) - + start_time = time.time() async with websockets.connect( auth_url, ping_interval=30, @@ -318,10 +447,7 @@ class StreamTTSPerformanceTester: ) as ws: # 构造请求 request = self._build_xunfei_request(app_id, text, voice) - # 发送请求后立即计时,确保准确测量从发送文本到接收首块的时间 await ws.send(json.dumps(request)) - start_time = time.time() - # 等待第一个音频数据块 first_audio_received = False while not first_audio_received: @@ -329,14 +455,14 @@ class StreamTTSPerformanceTester: data = json.loads(msg) header = data.get("header", {}) code = header.get("code") - + if code != 0: message = header.get("message", "未知错误") raise Exception(f"合成失败: {code} - {message}") - + payload = data.get("payload", {}) audio_payload = payload.get("audio", {}) - + if audio_payload: status = audio_payload.get("status", 0) audio_data = audio_payload.get("audio", "") @@ -344,10 +470,12 @@ class StreamTTSPerformanceTester: # 收到第一个音频数据块 latency = time.time() - start_time latencies.append(latency) + print(f"[讯飞TTS] 第{i+1}次 首词延迟: {latency:.3f}s") first_audio_received = True break except Exception as e: - latencies.append(0) + print(f"[讯飞TTS] 第{i+1}次测试失败: {str(e)}") + latencies.append(None) return self._calculate_result("讯飞TTS", latencies, test_count) @@ -431,8 +559,9 @@ class StreamTTSPerformanceTester: def _calculate_result(self, service_name, latencies, test_count): - """计算测试结果""" - valid_latencies = [l for l in latencies if l > 0] + """计算测试结果(正确处理None值,剔除失败测试)""" + # 剔除失败的测试(None值和<=0延迟),只统计有效延迟 + valid_latencies = [l for l in latencies if l is not None and l > 0] if valid_latencies: avg_latency = sum(valid_latencies) / len(valid_latencies) status = f"成功({len(valid_latencies)}/{test_count}次有效)" @@ -466,9 +595,10 @@ class StreamTTSPerformanceTester: ] print(tabulate(table_data, headers=["TTS服务", "首词延迟(秒)", "状态"], tablefmt="grid")) - print("\n测试说明:测量从发送请求到接收第一个音频数据块的时间,取多次测试平均值") + print("\n测试说明:测量从建立连接到接收第一个音频数据块的时间(包含握手、鉴权、发送文本),取多次测试平均值") + print("- 计时起点: 建立WebSocket/HTTP连接前(统一包含网络建连、握手、发送文本全流程)") print("- 超时控制: 单个请求最大等待时间为10秒") - print("- 错误处理: 无法连接和超时的列为网络错误") + print("- 错误处理: 失败的测试不计入平均值,只统计成功测试的延迟") print("- 排序规则: 按平均耗时从快到慢排序") @@ -494,7 +624,12 @@ class StreamTTSPerformanceTester: # 测试阿里云TTS result = await self.test_aliyun_tts(test_text, test_count) self.results.append(result) - + + # 测试阿里云百炼TTS + if self.config.get("TTS", {}).get("AliBLTTS"): + result = await self.test_alibl_tts(test_text, test_count) + self.results.append(result) + # 测试火山引擎TTS result = await self.test_doubao_tts(test_text, test_count) self.results.append(result) diff --git a/main/xiaozhi-server/plugins_func/functions/get_news_from_chinanews.py b/main/xiaozhi-server/plugins_func/functions/get_news_from_chinanews.py index e5ca0d1f..2267c83c 100644 --- a/main/xiaozhi-server/plugins_func/functions/get_news_from_chinanews.py +++ b/main/xiaozhi-server/plugins_func/functions/get_news_from_chinanews.py @@ -195,7 +195,7 @@ def get_news_from_chinanews( # 否则,获取新闻列表并随机选择一条 # 从配置中获取RSS URL - rss_config = conn.config["plugins"]["get_news_from_chinanews"] + rss_config = conn.config.get("plugins", {}).get("get_news_from_chinanews", {}) default_rss_url = rss_config.get( "default_rss_url", "https://www.chinanews.com.cn/rss/society.xml" ) diff --git a/main/xiaozhi-server/plugins_func/functions/get_news_from_newsnow.py b/main/xiaozhi-server/plugins_func/functions/get_news_from_newsnow.py index 1d60aefd..2bcd9193 100644 --- a/main/xiaozhi-server/plugins_func/functions/get_news_from_newsnow.py +++ b/main/xiaozhi-server/plugins_func/functions/get_news_from_newsnow.py @@ -120,10 +120,10 @@ def fetch_news_from_api(conn, source="thepaper"): """从API获取新闻列表""" try: api_url = f"https://newsnow.busiyi.world/api/s?id={source}" - if conn.config["plugins"].get("get_news_from_newsnow") and conn.config[ - "plugins" - ]["get_news_from_newsnow"].get("url"): - api_url = conn.config["plugins"]["get_news_from_newsnow"]["url"] + source + + news_config = conn.config.get("plugins", {}).get("get_news_from_newsnow", {}) + if news_config.get("url"): + api_url = news_config["url"] + source headers = {"User-Agent": "Mozilla/5.0"} response = requests.get(api_url, headers=headers, timeout=10) diff --git a/main/xiaozhi-server/plugins_func/functions/get_weather.py b/main/xiaozhi-server/plugins_func/functions/get_weather.py index 38770a3f..e95a40d8 100644 --- a/main/xiaozhi-server/plugins_func/functions/get_weather.py +++ b/main/xiaozhi-server/plugins_func/functions/get_weather.py @@ -158,13 +158,10 @@ def parse_weather_info(soup): def get_weather(conn, location: str = None, lang: str = "zh_CN"): from core.utils.cache.manager import cache_manager, CacheType - api_host = conn.config["plugins"]["get_weather"].get( - "api_host", "mj7p3y7naa.re.qweatherapi.com" - ) - api_key = conn.config["plugins"]["get_weather"].get( - "api_key", "a861d0d5e7bf4ee1a83d9a9e4f96d4da" - ) - default_location = conn.config["plugins"]["get_weather"]["default_location"] + weather_config = conn.config.get("plugins", {}).get("get_weather", {}) + api_host = weather_config.get("api_host", "mj7p3y7naa.re.qweatherapi.com") + api_key = weather_config.get("api_key", "a861d0d5e7bf4ee1a83d9a9e4f96d4da") + default_location = weather_config.get("default_location", "广州") client_ip = conn.client_ip # 优先使用用户提供的location参数 diff --git a/main/xiaozhi-server/plugins_func/functions/hass_init.py b/main/xiaozhi-server/plugins_func/functions/hass_init.py index 11cbb7a0..dadb190b 100644 --- a/main/xiaozhi-server/plugins_func/functions/hass_init.py +++ b/main/xiaozhi-server/plugins_func/functions/hass_init.py @@ -11,15 +11,17 @@ def append_devices_to_prompt(conn): "functions", [] ) + # 安全地获取插件配置 + plugins_config = conn.config.get("plugins", {}) config_source = ( "home_assistant" - if conn.config["plugins"].get("home_assistant") + if plugins_config.get("home_assistant") else "hass_get_state" ) if "hass_get_state" in funcs or "hass_set_state" in funcs: prompt = "\n下面是我家智能设备列表(位置,设备名,entity_id),可以通过homeassistant控制\n" - deviceStr = conn.config["plugins"].get(config_source, {}).get("devices", "") + deviceStr = plugins_config.get(config_source, {}).get("devices", "") conn.prompt += prompt + deviceStr + "\n" # 更新提示词 conn.dialogue.update_system_message(conn.prompt) @@ -30,17 +32,17 @@ def initialize_hass_handler(conn): if not conn.load_function_plugin: return ha_config + # 安全地获取插件配置 + plugins_config = conn.config.get("plugins", {}) # 确定配置来源 config_source = ( - "home_assistant" - if conn.config["plugins"].get("home_assistant") - else "hass_get_state" + "home_assistant" if plugins_config.get("home_assistant") else "hass_get_state" ) - if not conn.config["plugins"].get(config_source): + if not plugins_config.get(config_source): return ha_config # 统一获取配置 - plugin_config = conn.config["plugins"][config_source] + plugin_config = plugins_config[config_source] ha_config["base_url"] = plugin_config.get("base_url") ha_config["api_key"] = plugin_config.get("api_key") diff --git a/main/xiaozhi-server/plugins_func/functions/play_music.py b/main/xiaozhi-server/plugins_func/functions/play_music.py index 2cbc4018..be4cf618 100644 --- a/main/xiaozhi-server/plugins_func/functions/play_music.py +++ b/main/xiaozhi-server/plugins_func/functions/play_music.py @@ -118,8 +118,9 @@ def get_music_files(music_dir, music_ext): def initialize_music_handler(conn): global MUSIC_CACHE if MUSIC_CACHE == {}: - if "play_music" in conn.config["plugins"]: - MUSIC_CACHE["music_config"] = conn.config["plugins"]["play_music"] + plugins_config = conn.config.get("plugins", {}) + if "play_music" in plugins_config: + MUSIC_CACHE["music_config"] = plugins_config["play_music"] MUSIC_CACHE["music_dir"] = os.path.abspath( MUSIC_CACHE["music_config"].get("music_dir", "./music") # 默认路径修改 ) diff --git a/main/xiaozhi-server/plugins_func/functions/search_from_ragflow.py b/main/xiaozhi-server/plugins_func/functions/search_from_ragflow.py index d585a762..ec6ac426 100644 --- a/main/xiaozhi-server/plugins_func/functions/search_from_ragflow.py +++ b/main/xiaozhi-server/plugins_func/functions/search_from_ragflow.py @@ -32,9 +32,10 @@ def search_from_ragflow(conn, question=None): else: question = str(question) if question is not None else "" - base_url = conn.config["plugins"]["search_from_ragflow"].get("base_url", "") - api_key = conn.config["plugins"]["search_from_ragflow"].get("api_key", "") - dataset_ids = conn.config["plugins"]["search_from_ragflow"].get("dataset_ids", []) + ragflow_config = conn.config.get("plugins", {}).get("search_from_ragflow", {}) + base_url = ragflow_config.get("base_url", "") + api_key = ragflow_config.get("api_key", "") + dataset_ids = ragflow_config.get("dataset_ids", []) url = base_url + "/api/v1/retrieval" headers = {"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"} diff --git a/main/xiaozhi-server/requirements.txt b/main/xiaozhi-server/requirements.txt index a5b791ff..919c0f6d 100644 --- a/main/xiaozhi-server/requirements.txt +++ b/main/xiaozhi-server/requirements.txt @@ -1,3 +1,4 @@ +# -*- coding:utf-8 -*- #--------- 本项目推荐环境是python3.10,以下暂时不推荐升级的依赖 torch==2.2.2 torchaudio==2.2.2 @@ -10,9 +11,9 @@ silero_vad==6.1.0 opuslib_next==1.1.5 pydub==0.25.1 funasr==1.2.7 -openai==2.7.1 +openai==2.8.1 google-generativeai==0.8.5 -edge_tts==7.2.3 +edge_tts==7.2.6 httpx==0.28.1 aiohttp==3.13.2 aiohttp_cors==0.8.1 @@ -24,11 +25,11 @@ cozepy==0.20.0 mem0ai==1.0.0 bs4==0.0.2 modelscope==1.23.2 -sherpa_onnx==1.12.15 +sherpa_onnx==1.12.17 mcp==1.20.0 cnlunar==0.2.0 PySocks==1.7.1 -dashscope==1.24.6 +dashscope==1.25.2 baidu-aip==4.16.13 chardet==5.2.0 aioconsole==0.8.2 @@ -38,4 +39,4 @@ PyJWT==2.10.1 psutil==7.0.0 portalocker==3.2.0 Jinja2==3.1.6 -vosk==0.3.44 \ No newline at end of file +vosk==0.3.45 \ No newline at end of file diff --git a/main/xiaozhi-server/test/css/test_page.css b/main/xiaozhi-server/test/css/test_page.css index c0564191..b197c60d 100644 --- a/main/xiaozhi-server/test/css/test_page.css +++ b/main/xiaozhi-server/test/css/test_page.css @@ -267,6 +267,7 @@ span.connection-status.llm-emoji { .llm-emoji .status { font-size: 14px !important; padding: 8px 20px !important; + line-height: 1.2 !important; } .emoji-large { @@ -393,7 +394,7 @@ span.connection-status.llm-emoji { /* ==================== 会话记录和日志 ==================== */ .flex-container { display: flex; - margin-top: 10px; + margin-top: 20px; background-color: #f9fafb; } diff --git a/main/xiaozhi-server/test/js/core/audio/player.js b/main/xiaozhi-server/test/js/core/audio/player.js index 459fbb54..bc89d58a 100644 --- a/main/xiaozhi-server/test/js/core/audio/player.js +++ b/main/xiaozhi-server/test/js/core/audio/player.js @@ -249,6 +249,41 @@ export class AudioPlayer { this.playBufferedAudio(); this.startAudioBuffering(); } + + // 获取音频包统计信息 + getAudioStats() { + if (!this.streamingContext) { + return { + pendingDecode: 0, + pendingPlay: 0, + totalPending: 0 + }; + } + + const pendingDecode = this.streamingContext.getPendingDecodeCount(); + const pendingPlay = this.streamingContext.getPendingPlayCount(); + + return { + pendingDecode, // 待解码包数 + pendingPlay, // 待播放包数 + totalPending: pendingDecode + pendingPlay // 总待处理包数 + }; + } + + // 清空所有音频缓冲并停止播放 + clearAllAudio() { + log('AudioPlayer: 清空所有音频', 'info'); + + // 清空接收队列(使用clear方法保持对象引用) + this.queue.clear(); + + // 清空流上下文的所有缓冲 + if (this.streamingContext) { + this.streamingContext.clearAllBuffers(); + } + + log('AudioPlayer: 音频已清空', 'success'); + } } // 创建单例 diff --git a/main/xiaozhi-server/test/js/core/audio/recorder.js b/main/xiaozhi-server/test/js/core/audio/recorder.js index 9f94f895..3fe785f0 100644 --- a/main/xiaozhi-server/test/js/core/audio/recorder.js +++ b/main/xiaozhi-server/test/js/core/audio/recorder.js @@ -240,6 +240,24 @@ export class AudioRecorder { if (this.isRecording) return false; try { + // 检查是否有WebSocketHandler实例 + const { getWebSocketHandler } = await import('../network/websocket.js'); + const wsHandler = getWebSocketHandler(); + + // 如果机器正在说话,发送打断消息 + if (wsHandler && wsHandler.isRemoteSpeaking && wsHandler.currentSessionId) { + const abortMessage = { + session_id: wsHandler.currentSessionId, + type: 'abort', + reason: 'wake_word_detected' + }; + + if (this.websocket && this.websocket.readyState === WebSocket.OPEN) { + this.websocket.send(JSON.stringify(abortMessage)); + log('发送打断消息', 'info'); + } + } + if (!this.initEncoder()) { log('无法启动录音: Opus编码器初始化失败', 'error'); return false; diff --git a/main/xiaozhi-server/test/js/core/audio/stream-context.js b/main/xiaozhi-server/test/js/core/audio/stream-context.js index 9c722505..05b8c1db 100644 --- a/main/xiaozhi-server/test/js/core/audio/stream-context.js +++ b/main/xiaozhi-server/test/js/core/audio/stream-context.js @@ -22,6 +22,7 @@ export class StreamingContext { this.source = null; // 当前音频源 this.totalSamples = 0; // 累积的总样本数 this.lastPlayTime = 0; // 上次播放的时间戳 + this.scheduledEndTime = 0; // 已调度音频的结束时间 } // 缓存音频数组 @@ -31,17 +32,19 @@ export class StreamingContext { // 获取需要处理缓存队列,单线程:在audioBufferQueue一直更新的状态下不会出现安全问题 async getPendingAudioBufferQueue() { - // 原子交换 + 清空 - [this.pendingAudioBufferQueue, this.audioBufferQueue] = [await this.audioBufferQueue.dequeue(), new BlockingQueue()]; + // 等待数据到达并获取 + const data = await this.audioBufferQueue.dequeue(); + // 赋值给待处理队列 + this.pendingAudioBufferQueue = data; } // 获取正在播放已解码的PCM队列,单线程:在activeQueue一直更新的状态下不会出现安全问题 async getQueue(minSamples) { - let TepArray = []; const num = minSamples - this.queue.length > 0 ? minSamples - this.queue.length : 1; - // 原子交换 + 清空 - [TepArray, this.activeQueue] = [await this.activeQueue.dequeue(num), new BlockingQueue()]; - this.queue.push(...TepArray); + + // 等待数据并获取 + const tempArray = await this.activeQueue.dequeue(num); + this.queue.push(...tempArray); } // 将Int16音频数据转换为Float32音频数据 @@ -54,6 +57,57 @@ export class StreamingContext { return float32Data; } + // 获取待解码包数 + getPendingDecodeCount() { + return this.audioBufferQueue.length + this.pendingAudioBufferQueue.length; + } + + // 获取待播放样本数(转换为包数,每包960样本) + getPendingPlayCount() { + // 计算已在队列中的样本 + const queuedSamples = this.activeQueue.length + this.queue.length; + + // 计算已调度但未播放的样本(在Web Audio缓冲区中) + let scheduledSamples = 0; + if (this.playing && this.scheduledEndTime) { + const currentTime = this.audioContext.currentTime; + const remainingTime = Math.max(0, this.scheduledEndTime - currentTime); + scheduledSamples = Math.floor(remainingTime * this.sampleRate); + } + + const totalSamples = queuedSamples + scheduledSamples; + return Math.ceil(totalSamples / 960); + } + + // 清空所有音频缓冲 + clearAllBuffers() { + log('清空所有音频缓冲', 'info'); + + // 清空所有队列(使用clear方法保持对象引用) + this.audioBufferQueue.clear(); + this.pendingAudioBufferQueue = []; + this.activeQueue.clear(); + this.queue = []; + + // 停止当前播放的音频源 + if (this.source) { + try { + this.source.stop(); + this.source.disconnect(); + } catch (e) { + // 忽略已经停止的错误 + } + this.source = null; + } + + // 重置状态 + this.playing = false; + this.scheduledEndTime = this.audioContext.currentTime; + this.totalSamples = 0; + + log('音频缓冲已清空', 'success'); + } + // 将Opus数据解码为PCM async decodeOpusFrames() { if (!this.opusDecoder) { @@ -97,7 +151,7 @@ export class StreamingContext { // 开始播放音频 async startPlaying() { - let scheduledEndTime = this.audioContext.currentTime; // 跟踪已调度音频的结束时间 + this.scheduledEndTime = this.audioContext.currentTime; // 跟踪已调度音频的结束时间 while (true) { // 初始缓冲:等待足够的样本再开始播放 @@ -126,7 +180,7 @@ export class StreamingContext { // 精确调度播放时间 const currentTime = this.audioContext.currentTime; - const startTime = Math.max(scheduledEndTime, currentTime); + const startTime = Math.max(this.scheduledEndTime, currentTime); // 直接连接到输出 this.source.connect(this.audioContext.destination); @@ -136,7 +190,7 @@ export class StreamingContext { // 更新下一个音频块的调度时间 const duration = audioBuffer.duration; - scheduledEndTime = startTime + duration; + this.scheduledEndTime = startTime + duration; this.lastPlayTime = startTime; // 如果队列中数据不足,等待新数据 diff --git a/main/xiaozhi-server/test/js/core/network/websocket.js b/main/xiaozhi-server/test/js/core/network/websocket.js index b62c710b..a605e6f1 100644 --- a/main/xiaozhi-server/test/js/core/network/websocket.js +++ b/main/xiaozhi-server/test/js/core/network/websocket.js @@ -123,7 +123,12 @@ export class WebSocketHandler { } else if (message.state === 'sentence_end') { log(`语音段结束: ${message.text}`, 'info'); } else if (message.state === 'stop') { - log('服务器语音传输结束', 'info'); + log('服务器语音传输结束,清空所有音频缓冲', 'info'); + + // 清空所有音频缓冲并停止播放 + const audioPlayer = getAudioPlayer(); + audioPlayer.clearAllAudio(); + this.isRemoteSpeaking = false; if (this.onRecordButtonStateChange) { this.onRecordButtonStateChange(false); diff --git a/main/xiaozhi-server/test/js/ui/controller.js b/main/xiaozhi-server/test/js/ui/controller.js index 8fc9dd31..3e4f6563 100644 --- a/main/xiaozhi-server/test/js/ui/controller.js +++ b/main/xiaozhi-server/test/js/ui/controller.js @@ -2,6 +2,7 @@ import { loadConfig, saveConfig } from '../config/manager.js'; import { getAudioRecorder } from '../core/audio/recorder.js'; import { getWebSocketHandler } from '../core/network/websocket.js'; +import { getAudioPlayer } from '../core/audio/player.js'; // UI控制器类 export class UIController { @@ -9,6 +10,7 @@ export class UIController { this.isEditing = false; this.visualizerCanvas = null; this.visualizerContext = null; + this.audioStatsTimer = null; } // 初始化 @@ -18,6 +20,7 @@ export class UIController { this.initVisualizer(); this.initEventListeners(); + this.startAudioStatsMonitor(); loadConfig(); } @@ -86,17 +89,20 @@ export class UIController { const sessionStatus = document.getElementById('sessionStatus'); if (!sessionStatus) return; + // 保留背景元素 + const bgHtml = ''; + if (isSpeaking === null) { // 离线状态 - sessionStatus.innerHTML = '😶 小智离线中'; + sessionStatus.innerHTML = bgHtml + '😶 小智离线中'; sessionStatus.className = 'status offline'; } else if (isSpeaking) { // 说话中 - sessionStatus.innerHTML = '😶 小智说话中'; + sessionStatus.innerHTML = bgHtml + '😶 小智说话中'; sessionStatus.className = 'status speaking'; } else { // 聆听中 - sessionStatus.innerHTML = '😶 小智聆听中'; + sessionStatus.innerHTML = bgHtml + '😶 小智聆听中'; sessionStatus.className = 'status listening'; } } @@ -110,8 +116,72 @@ export class UIController { let currentText = sessionStatus.textContent; // 移除现有的表情符号 currentText = currentText.replace(/[\u{1F300}-\u{1F9FF}]|[\u{2600}-\u{26FF}]|[\u{2700}-\u{27BF}]/gu, '').trim(); + + // 保留背景元素 + const bgHtml = ''; + // 使用 innerHTML 添加带样式的表情 - sessionStatus.innerHTML = `${emoji} ${currentText}`; + sessionStatus.innerHTML = bgHtml + `${emoji} ${currentText}`; + } + + // 更新音频统计信息 + updateAudioStats() { + const audioPlayer = getAudioPlayer(); + const stats = audioPlayer.getAudioStats(); + + const sessionStatus = document.getElementById('sessionStatus'); + const sessionStatusBg = document.getElementById('sessionStatusBg'); + + // 只在说话状态下显示背景进度 + if (sessionStatus && sessionStatus.classList.contains('speaking') && sessionStatusBg) { + if (stats.pendingPlay > 0) { + // 计算进度:5包=50%,10包及以上=100% + let percentage; + if (stats.pendingPlay >= 10) { + percentage = 100; + } else { + percentage = (stats.pendingPlay / 10) * 100; + } + + sessionStatusBg.style.width = `${percentage}%`; + + // 根据缓冲量改变背景颜色 + if (stats.pendingPlay < 5) { + // 缓冲不足:橙红色半透明 + sessionStatusBg.style.background = 'linear-gradient(90deg, rgba(255, 152, 0, 0.25), rgba(255, 87, 34, 0.25))'; + } else if (stats.pendingPlay < 10) { + // 一般:黄绿色半透明 + sessionStatusBg.style.background = 'linear-gradient(90deg, rgba(205, 220, 57, 0.25), rgba(76, 175, 80, 0.25))'; + } else { + // 充足:绿蓝色半透明 + sessionStatusBg.style.background = 'linear-gradient(90deg, rgba(76, 175, 80, 0.25), rgba(33, 150, 243, 0.25))'; + } + } else { + // 没有缓冲,隐藏背景 + sessionStatusBg.style.width = '0%'; + } + } else { + // 非说话状态,隐藏背景 + if (sessionStatusBg) { + sessionStatusBg.style.width = '0%'; + } + } + } + + // 启动音频统计监控 + startAudioStatsMonitor() { + // 每100ms更新一次音频统计 + this.audioStatsTimer = setInterval(() => { + this.updateAudioStats(); + }, 100); + } + + // 停止音频统计监控 + stopAudioStatsMonitor() { + if (this.audioStatsTimer) { + clearInterval(this.audioStatsTimer); + this.audioStatsTimer = null; + } } // 绘制音频可视化效果 diff --git a/main/xiaozhi-server/test/js/utils/blocking-queue.js b/main/xiaozhi-server/test/js/utils/blocking-queue.js index 31a43872..738a3e75 100644 --- a/main/xiaozhi-server/test/js/utils/blocking-queue.js +++ b/main/xiaozhi-server/test/js/utils/blocking-queue.js @@ -95,4 +95,9 @@ export default class BlockingQueue { get length() { return this.#items.length; } + + /* 清空队列(保持对象引用,不影响等待者) */ + clear() { + this.#items.length = 0; + } } \ No newline at end of file diff --git a/main/xiaozhi-server/test/test_page.html b/main/xiaozhi-server/test/test_page.html index 97af77ac..f3f583c9 100644 --- a/main/xiaozhi-server/test/test_page.html +++ b/main/xiaozhi-server/test/test_page.html @@ -91,6 +91,7 @@
+
@@ -113,9 +114,12 @@
-
+
- 😶 小智离线中 + + + 😶 小智离线中 +