diff --git a/.github/workflows/docker-image.yml b/.github/workflows/docker-image.yml index e54cdd4e..aefa77d0 100644 --- a/.github/workflows/docker-image.yml +++ b/.github/workflows/docker-image.yml @@ -31,6 +31,9 @@ jobs: - name: Set up Docker Buildx uses: docker/setup-buildx-action@v3 + with: + driver-opts: | + network=host - name: Login to GitHub Container Registry uses: docker/login-action@v3 @@ -60,6 +63,10 @@ jobs: tags: | ${{ env.IS_VERSION == 'true' && format('ghcr.io/{0}:server_{1},ghcr.io/{0}:server_latest', github.repository, env.VERSION) || format('ghcr.io/{0}:server_latest', github.repository) }} platforms: linux/amd64,linux/arm64 + cache-from: type=gha + cache-to: type=gha,mode=max + build-args: | + BUILDKIT_PROGRESS=plain # 构建 manager-api 镜像 - name: Build and push manager-web @@ -70,4 +77,8 @@ jobs: push: true tags: | ${{ env.IS_VERSION == 'true' && format('ghcr.io/{0}:web_{1},ghcr.io/{0}:web_latest', github.repository, env.VERSION) || format('ghcr.io/{0}:web_latest', github.repository) }} - platforms: linux/amd64,linux/arm64 \ No newline at end of file + platforms: linux/amd64,linux/arm64 + cache-from: type=gha + cache-to: type=gha,mode=max + build-args: | + BUILDKIT_PROGRESS=plain \ No newline at end of file diff --git a/Dockerfile-server b/Dockerfile-server index 6e87f7c0..a12fbb11 100644 --- a/Dockerfile-server +++ b/Dockerfile-server @@ -3,10 +3,17 @@ FROM python:3.10-slim AS builder WORKDIR /app +# 配置pip使用国内镜像源(阿里云)并设置超时和重试 +RUN pip config set global.index-url https://mirrors.aliyun.com/pypi/simple/ && \ + pip config set global.trusted-host mirrors.aliyun.com && \ + pip config set global.timeout 120 && \ + pip config set install.retries 5 + COPY main/xiaozhi-server/requirements.txt . -# 安装Python依赖 -RUN pip install --no-cache-dir -r requirements.txt +# 安装Python依赖,使用并行下载 +RUN pip install --no-cache-dir --upgrade pip setuptools wheel && \ + pip install --no-cache-dir -r requirements.txt --default-timeout=120 --retries 5 # 第二阶段:生产镜像 FROM python:3.10-slim diff --git a/Dockerfile-web b/Dockerfile-web index 35c25ed2..a347a88e 100644 --- a/Dockerfile-web +++ b/Dockerfile-web @@ -18,12 +18,12 @@ FROM bellsoft/liberica-runtime-container:jre-21-glibc # 安装Nginx和字体库 RUN apk update && \ - apk add --no-cache nginx bash && \ - apk add --no-cache fontconfig ttf-dejavu msttcorefonts-installer && \ + apk add --no-cache nginx bash fontconfig ttf-dejavu && \ + apk add --no-cache --repository=http://dl-cdn.alpinelinux.org/alpine/edge/testing/ msttcorefonts-installer || true && \ rm -rf /var/cache/apk/* # 更新字体缓存 -RUN printf 'YES\n' | update-ms-fonts && fc-cache -f -v +RUN (printf 'YES\n' | update-ms-fonts || true) && fc-cache -f -v # 配置Nginx COPY docs/docker/nginx.conf /etc/nginx/nginx.conf diff --git a/README.md b/README.md index 8480e2be..862e989a 100644 --- a/README.md +++ b/README.md @@ -94,7 +94,7 @@ Spearheaded by Professor Siyuan Liu's Team (South China University of Technology - + 自定义音色 @@ -283,6 +283,8 @@ Websocket接口地址: wss://2662r3426b.vicp.fun/xiaozhi/v1/ | dify 接口调用 | Dify | - | | fastgpt 接口调用 | Fastgpt | - | | coze 接口调用 | Coze | - | +| xinference 接口调用 | Xinference | - | +| homeassistant 接口调用 | HomeAssistant | - | 实际上,任何支持 openai 接口调用的 LLM 均可接入使用。 @@ -302,8 +304,8 @@ Websocket接口地址: wss://2662r3426b.vicp.fun/xiaozhi/v1/ | 使用方式 | 支持平台 | 免费平台 | |:---:|:---:|:---:| -| 接口调用 | EdgeTTS、火山引擎豆包TTS、腾讯云、阿里云TTS、CosyVoiceSiliconflow、TTS302AI、CozeCnTTS、GizwitsTTS、ACGNTTS、OpenAITTS、灵犀流式TTS | 灵犀流式TTS、EdgeTTS、CosyVoiceSiliconflow(部分) | -| 本地服务 | FishSpeech、GPT_SOVITS_V2、GPT_SOVITS_V3、MinimaxTTS | FishSpeech、GPT_SOVITS_V2、GPT_SOVITS_V3、MinimaxTTS | +| 接口调用 | EdgeTTS、火山引擎豆包TTS、腾讯云、阿里云TTS、阿里云流式TTS、CosyVoiceSiliconflow、TTS302AI、CozeCnTTS、GizwitsTTS、ACGNTTS、OpenAITTS、灵犀流式TTS、MinimaxTTS、火山双流式TTS | 灵犀流式TTS、EdgeTTS、CosyVoiceSiliconflow(部分) | +| 本地服务 | FishSpeech、GPT_SOVITS_V2、GPT_SOVITS_V3、Index-TTS、PaddleSpeech | Index-TTS、PaddleSpeech、FishSpeech、GPT_SOVITS_V2、GPT_SOVITS_V3 | --- @@ -320,7 +322,7 @@ Websocket接口地址: wss://2662r3426b.vicp.fun/xiaozhi/v1/ | 使用方式 | 支持平台 | 免费平台 | |:---:|:---:|:---:| | 本地使用 | FunASR、SherpaASR | FunASR、SherpaASR | -| 接口调用 | DoubaoASR、FunASRServer、TencentASR、AliyunASR | FunASRServer | +| 接口调用 | DoubaoASR、Doubao流式ASR、FunASRServer、TencentASR、AliyunASR、Aliyun流式ASR、百度ASR、OpenAI ASR | FunASRServer | --- @@ -338,6 +340,7 @@ Websocket接口地址: wss://2662r3426b.vicp.fun/xiaozhi/v1/ |:------:|:---------------:|:----:|:---------:|:--:| | Memory | mem0ai | 接口调用 | 1000次/月额度 | | | Memory | mem_local_short | 本地总结 | 免费 | | +| Memory | nomem | 无记忆模式 | 免费 | | --- @@ -347,6 +350,7 @@ Websocket接口地址: wss://2662r3426b.vicp.fun/xiaozhi/v1/ |:------:|:-------------:|:----:|:-------:|:---------------------:| | Intent | intent_llm | 接口调用 | 根据LLM收费 | 通过大模型识别意图,通用性强 | | Intent | function_call | 接口调用 | 根据LLM收费 | 通过大模型函数调用完成意图,速度快,效果好 | +| Intent | nointent | 无意图模式 | 免费 | 不进行意图识别,直接返回对话结果 | --- diff --git a/README_en.md b/README_en.md index 85086968..5432eae8 100644 --- a/README_en.md +++ b/README_en.md @@ -93,7 +93,7 @@ Want to see the usage effects? Click the videos below 🎥 - + Custom voice timbre diff --git a/docs/Deployment_all.md b/docs/Deployment_all.md index 135b7cf5..96b3d75c 100644 --- a/docs/Deployment_all.md +++ b/docs/Deployment_all.md @@ -476,7 +476,8 @@ ws://你电脑局域网的ip:8000/xiaozhi/v1/ 4、[如何部署MCP接入点](./mcp-endpoint-enable.md)
5、[如何接入MCP接入点](./mcp-endpoint-integration.md)
6、[如何开启声纹识别](./voiceprint-integration.md)
-10、[新闻插件源配置指南](./newsnow_plugin_config.md)
+7、[新闻插件源配置指南](./newsnow_plugin_config.md)
+8、[天气插件使用指南](./weather-integration.md)
## 语音克隆、本地语音部署相关教程 1、[如何部署集成index-tts本地语音](./index-stream-integration.md)
2、[如何部署集成fish-speech本地语音](./fish-speech-integration.md)
diff --git a/docs/images/demo6.png b/docs/images/demo6.png index 18d8f5f3..dc4edbeb 100644 Binary files a/docs/images/demo6.png and b/docs/images/demo6.png differ diff --git a/docs/paddlespeech-deploy.md b/docs/paddlespeech-deploy.md index 6584efff..17030fab 100644 --- a/docs/paddlespeech-deploy.md +++ b/docs/paddlespeech-deploy.md @@ -75,7 +75,7 @@ TTS: sample_rate: 24000 # 采样率 [websocket默认24000,http默认0 自动选择] speed: 1.0 # 语速,1.0 表示正常语速,>1 表示加快,<1 表示减慢 volume: 1.0 # 音量,1.0 表示正常音量,>1 表示增大,<1 表示减小 - save_path: ./streaming_tts.wav # 服务器生成的语音文件保存路径 + save_path: # 保存路径 ``` ### 3.启动xiaozhi服务 ```py diff --git a/docs/weather-integration.md b/docs/weather-integration.md new file mode 100644 index 00000000..3b6ca2b6 --- /dev/null +++ b/docs/weather-integration.md @@ -0,0 +1,64 @@ +# 天气插件使用指南 + +## 概述 + +天气插件 `get_weather` 是小智ESP32语音助手的核心功能之一,支持通过语音查询全国各地的天气信息。插件基于和风天气API,提供实时天气和7天天气预报功能。 + +## API Key 申请指南 + +### 1. 注册和风天气账号 + +1. 访问 [和风天气控制台](https://console.qweather.com/) +2. 注册账号并完成邮箱验证 +3. 登录控制台 + +### 2. 创建应用获取API Key + +1. 进入控制台后,点击右侧["项目管理"](https://console.qweather.com/project?lang=zh) → "创建项目" +2. 填写项目信息: + - **项目名称**:如"小智语音助手" +3. 点击保存 +4. 项目创建完成后,在该项目中点击"创建凭据" +5. 填写凭据信息: + - **凭据名称**:如"小智语音助手" + - **身份认证方式**:选择"API Key" +6. 点击保存 +7. 在凭据中复制`API Key`,这是第一个关键的配置信息 + +### 3. 获取API Host + +1. 在控制台中点击["设置"](https://console.qweather.com/setting?lang=zh) → "API Host" +2. 查看分配给你的专属`API Host`地址,这个是第二个关键的配置信息 + +以上操作,会得到两个重要的配置信息:`API Key`和`API Host` + +## 配置方式(任选一种) + +### 方式1. 如果你使用了智控台部署(推荐) + +1. 登录智控台 +2. 进入"角色配置"页面 +3. 选择要配置的智能体 +4. 点击"编辑功能"按钮 +5. 在右侧参数配置区域找到"天气查询"插件 +6. 勾选"天气查询" +7. 将复制过来的第一个关键配置`API Key`,填入到`天气插件 API 密钥`里 +8. 将复制过来的第二个关键配置`API Host`,填入到`开发者 API Host`里 +9. 保存配置,再保存智能体配置 + +### 方式2. 如果你只是单模块xiaozhi-server部署 + +在 `data/.config.yaml` 中配置: + +1. 将复制过来的第一个关键配置`API Key`,填入到`api_key`里 +2. 将复制过来的第二个关键配置`API Host`,填入到`api_host`里 +3. 将你所在的城市填入到`default_location`里,例如`广州` + +```yaml +plugins: + get_weather: + api_key: "你的和风天气API密钥" + api_host: "你的和风天气API主机地址" + default_location: "你的默认查询城市" +``` + 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 bf747600..afd7addf 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 @@ -237,7 +237,7 @@ public interface Constant { /** * 版本号 */ - public static final String VERSION = "0.7.5"; + public static final String VERSION = "0.7.7"; /** * 无效固件URL diff --git a/main/manager-mobile/src/pages/settings/index.vue b/main/manager-mobile/src/pages/settings/index.vue index 41d16c00..8b50f928 100644 --- a/main/manager-mobile/src/pages/settings/index.vue +++ b/main/manager-mobile/src/pages/settings/index.vue @@ -56,11 +56,11 @@ function getCacheInfo() { // 验证URL格式 function validateUrl() { urlError.value = '' - + if (!baseUrlInput.value) { return } - + if (!/^https?:\/\/.+\/xiaozhi$/.test(baseUrlInput.value)) { urlError.value = '请输入有效的服务端地址(以 http 或 https 开头,并以 /xiaozhi 结尾)' } @@ -70,7 +70,7 @@ function validateUrl() { async function testServerBaseUrl() { // 先清除错误信息 urlError.value = '' - + if (!baseUrlInput.value || !/^https?:\/\/.+\/xiaozhi$/.test(baseUrlInput.value)) { return false } @@ -113,20 +113,20 @@ async function saveServerBaseUrl() { clearAllCacheAfterUrlChange() uni.showModal({ - title: '重启应用', - content: '服务端地址已保存并清空缓存,是否立即重启生效?', - confirmText: '立即重启', - cancelText: '稍后', - success: (res) => { - if (res.confirm) { - restartApp() - } - else { - toast.success('已保存,可稍后手动重启应用') - } - }, - }) - } + title: '重启应用', + content: '服务端地址已保存并清空缓存,是否立即重启生效?', + confirmText: '立即重启', + cancelText: '稍后', + success: (res) => { + if (res.confirm) { + restartApp() + } + else { + toast.success('已保存,可稍后手动重启应用') + } + }, + }) +} // 重置为 env 默认 function resetServerBaseUrl() { @@ -222,7 +222,7 @@ function showAbout() { title: `关于${import.meta.env.VITE_APP_TITLE}`, content: `${import.meta.env.VITE_APP_TITLE}\n\n基于 Vue.js 3 + uni-app 构建的跨平台移动端管理应用,为小智ESP32智能硬件提供设备管理、智能体配置等功能。\n\n© 2025 xiaozhi-esp32-server`, title: `关于小智智控台`, - content: `小智智控台\n\n基于 Vue.js 3 + uni-app 构建的跨平台移动端管理应用,为小智ESP32智能硬件提供设备管理、智能体配置等功能。\n\n© 2025 xiaozhi-esp32-server 0.7.5`, + content: `小智智控台\n\n基于 Vue.js 3 + uni-app 构建的跨平台移动端管理应用,为小智ESP32智能硬件提供设备管理、智能体配置等功能。\n\n© 2025 xiaozhi-esp32-server 0.7.7`, showCancel: false, confirmText: '确定', }) @@ -263,17 +263,10 @@ onMounted(async () => { - + input-class="text-[28rpx] text-[#232338]" @input="validateUrl" @blur="validateUrl" /> {{ urlError }} @@ -371,7 +364,7 @@ onMounted(async () => { - + diff --git a/main/manager-web/src/components/FunctionDialog.vue b/main/manager-web/src/components/FunctionDialog.vue index 76c69c8e..f8adfef7 100644 --- a/main/manager-web/src/components/FunctionDialog.vue +++ b/main/manager-web/src/components/FunctionDialog.vue @@ -698,15 +698,10 @@ export default { font-size: 14px; height: 36px; box-sizing: border-box; - background-color: #f5f5f5; -} -::v-deep .el-input__inner { - background-color: #f5f5f5; - padding-right: 80px; -} - -.url-input { + ::v-deep .el-input__inner { + background-color: #f5f5f5 !important; + } ::v-deep .el-input__suffix { right: 0; diff --git a/main/xiaozhi-server/config.yaml b/main/xiaozhi-server/config.yaml index 1e52f748..46e4253e 100644 --- a/main/xiaozhi-server/config.yaml +++ b/main/xiaozhi-server/config.yaml @@ -59,6 +59,12 @@ log: delete_audio: true # 没有语音输入多久后断开连接(秒),默认2分钟,即120秒 close_connection_no_voice_time: 120 +# TTS请求超时时间(秒) +tts_timeout: 10 +# 开启唤醒词加速 +enable_wakeup_words_response_cache: true +# 开场是否回复唤醒词 +enable_greeting: true # 说完话是否开启提示音 enable_stop_tts_notify: false # 说完话是否开启提示音,音效地址 @@ -911,7 +917,7 @@ TTS: sample_rate: 24000 # 采样率 [websocket默认24000,http默认0 自动选择] speed: 1.0 # 语速,1.0 表示正常语速,>1 表示加快,<1 表示减慢 volume: 1.0 # 音量,1.0 表示正常音量,>1 表示增大,<1 表示减小 - save_path: ./streaming_tts.wav # 服务器生成的语音文件保存路径 + save_path: # 保存路径 IndexStreamTTS: # 基于Index-TTS-vLLM项目的TTS接口服务 # 参照教程:https://github.com/Ksuriuri/index-tts-vllm/blob/master/README.md diff --git a/main/xiaozhi-server/config/logger.py b/main/xiaozhi-server/config/logger.py index cd4d9965..5c916ba5 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.7.5" +SERVER_VERSION = "0.7.7" _logger_initialized = False diff --git a/main/xiaozhi-server/core/handle/helloHandle.py b/main/xiaozhi-server/core/handle/helloHandle.py index de9587b5..75b9fcb2 100644 --- a/main/xiaozhi-server/core/handle/helloHandle.py +++ b/main/xiaozhi-server/core/handle/helloHandle.py @@ -1,5 +1,13 @@ +import time import json +import random import asyncio +from core.utils.dialogue import Message +from core.utils.util import audio_to_data +from core.providers.tts.dto.dto import SentenceType +from core.utils.wakeup_word import WakeupWordsConfig +from core.handle.sendAudioHandle import sendAudioMessage, send_stt_message +from core.utils.util import remove_punctuation_and_length, opus_datas_to_wav_bytes from core.providers.tools.device_mcp import ( MCPClient, send_mcp_initialize_message, @@ -8,6 +16,17 @@ from core.providers.tools.device_mcp import ( TAG = __name__ +WAKEUP_CONFIG = { + "refresh_time": 5, + "words": ["你好", "你好啊", "嘿,你好", "嗨"], +} + +# 创建全局的唤醒词配置管理器 +wakeup_words_config = WakeupWordsConfig() + +# 用于防止并发调用wakeupWordsResponse的锁 +_wakeup_response_lock = asyncio.Lock() + async def handleHelloMessage(conn, msg_json): """处理hello消息""" @@ -30,3 +49,103 @@ async def handleHelloMessage(conn, msg_json): asyncio.create_task(send_mcp_tools_list_request(conn)) await conn.websocket.send(json.dumps(conn.welcome_msg)) + + +async def checkWakeupWords(conn, text): + enable_wakeup_words_response_cache = conn.config[ + "enable_wakeup_words_response_cache" + ] + + # 等待tts初始化,最多等待3秒 + start_time = time.time() + while time.time() - start_time < 3: + if conn.tts: + break + await asyncio.sleep(0.1) + else: + return False + + if not enable_wakeup_words_response_cache: + return False + + _, filtered_text = remove_punctuation_and_length(text) + if filtered_text not in conn.config.get("wakeup_words"): + return False + + conn.just_woken_up = True + await send_stt_message(conn, text) + + # 获取当前音色 + voice = getattr(conn.tts, "voice", "default") + if not voice: + voice = "default" + + # 获取唤醒词回复配置 + response = wakeup_words_config.get_wakeup_response(voice) + if not response or not response.get("file_path"): + response = { + "voice": "default", + "file_path": "config/assets/wakeup_words.wav", + "time": 0, + "text": "哈啰啊,我是小智啦,声音好听的台湾女孩一枚,超开心认识你耶,最近在忙啥,别忘了给我来点有趣的料哦,我超爱听八卦的啦", + } + + # 获取音频数据 + opus_packets = audio_to_data(response.get("file_path")) + # 播放唤醒词回复 + conn.client_abort = False + + conn.logger.bind(tag=TAG).info(f"播放唤醒词回复: {response.get('text')}") + await sendAudioMessage(conn, SentenceType.FIRST, opus_packets, response.get("text")) + await sendAudioMessage(conn, SentenceType.LAST, [], None) + + # 补充对话 + conn.dialogue.put(Message(role="assistant", content=response.get("text"))) + + # 检查是否需要更新唤醒词回复 + if time.time() - response.get("time", 0) > WAKEUP_CONFIG["refresh_time"]: + if not _wakeup_response_lock.locked(): + asyncio.create_task(wakeupWordsResponse(conn)) + return True + + +async def wakeupWordsResponse(conn): + if not conn.tts or not conn.llm or not conn.llm.response_no_stream: + return + + try: + # 尝试获取锁,如果获取不到就返回 + if not await _wakeup_response_lock.acquire(): + return + + # 生成唤醒词回复 + wakeup_word = random.choice(WAKEUP_CONFIG["words"]) + question = ( + "此刻用户正在和你说```" + + wakeup_word + + "```。\n请你根据以上用户的内容进行20-30字回复。要符合系统设置的角色情感和态度,不要像机器人一样说话。\n" + + "请勿对这条内容本身进行任何解释和回应,请勿返回表情符号,仅返回对用户的内容的回复。" + ) + + result = conn.llm.response_no_stream(conn.config["prompt"], question) + if not result or len(result) == 0: + return + + # 生成TTS音频 + tts_result = await asyncio.to_thread(conn.tts.to_tts, result) + if not tts_result: + return + + # 获取当前音色 + voice = getattr(conn.tts, "voice", "default") + + wav_bytes = opus_datas_to_wav_bytes(tts_result, sample_rate=16000) + file_path = wakeup_words_config.generate_file_path(voice) + with open(file_path, "wb") as f: + f.write(wav_bytes) + # 更新配置 + wakeup_words_config.update_wakeup_response(voice, file_path, result) + finally: + # 确保在任何情况下都释放锁 + if _wakeup_response_lock.locked(): + _wakeup_response_lock.release() \ No newline at end of file diff --git a/main/xiaozhi-server/core/handle/intentHandler.py b/main/xiaozhi-server/core/handle/intentHandler.py index 58b2cee7..72424968 100644 --- a/main/xiaozhi-server/core/handle/intentHandler.py +++ b/main/xiaozhi-server/core/handle/intentHandler.py @@ -1,11 +1,12 @@ import json -import asyncio import uuid +import asyncio +from core.utils.dialogue import Message +from core.providers.tts.dto.dto import ContentType +from core.handle.helloHandle import checkWakeupWords +from plugins_func.register import Action, ActionResponse from core.handle.sendAudioHandle import send_stt_message from core.utils.util import remove_punctuation_and_length -from core.providers.tts.dto.dto import ContentType -from core.utils.dialogue import Message -from plugins_func.register import Action, ActionResponse from core.providers.tts.dto.dto import TTSMessageDTO, SentenceType TAG = __name__ @@ -23,10 +24,14 @@ async def handle_user_intent(conn, text): pass # 检查是否有明确的退出命令 - filtered_text = remove_punctuation_and_length(text)[1] + _, filtered_text = remove_punctuation_and_length(text) if await check_direct_exit(conn, filtered_text): return True + # 检查是否是唤醒词 + if await checkWakeupWords(conn, filtered_text): + return True + if conn.intent_type == "function_call": # 使用支持function calling的聊天方法,不再进行意图分析 return False diff --git a/main/xiaozhi-server/core/handle/receiveAudioHandle.py b/main/xiaozhi-server/core/handle/receiveAudioHandle.py index e5b96be2..8db50633 100644 --- a/main/xiaozhi-server/core/handle/receiveAudioHandle.py +++ b/main/xiaozhi-server/core/handle/receiveAudioHandle.py @@ -1,11 +1,11 @@ import time import json -from core.handle.sendAudioHandle import send_stt_message +import asyncio +from core.utils.util import audio_to_data +from core.handle.abortHandle import handleAbortMessage from core.handle.intentHandler import handle_user_intent from core.utils.output_counter import check_device_output_limit -from core.handle.abortHandle import handleAbortMessage -from core.handle.sendAudioHandle import SentenceType -from core.utils.util import audio_to_data_stream +from core.handle.sendAudioHandle import send_stt_message, SentenceType TAG = __name__ @@ -13,7 +13,14 @@ TAG = __name__ async def handleAudioMessage(conn, audio): # 当前片段是否有人说话 have_voice = conn.vad.is_vad(conn, audio) - + # 如果设备刚刚被唤醒,短暂忽略VAD检测 + if have_voice and hasattr(conn, "just_woken_up") and conn.just_woken_up: + have_voice = False + # 设置一个短暂延迟后恢复VAD检测 + conn.asr_audio.clear() + if not hasattr(conn, "vad_resume_task") or conn.vad_resume_task.done(): + conn.vad_resume_task = asyncio.create_task(resume_vad_detection(conn)) + return if have_voice: if conn.client_is_speaking: await handleAbortMessage(conn) @@ -22,6 +29,11 @@ async def handleAudioMessage(conn, audio): # 接收音频 await conn.asr.receive_audio(conn, audio, have_voice) +async def resume_vad_detection(conn): + # 等待2秒后恢复VAD检测 + await asyncio.sleep(1) + conn.just_woken_up = False + async def startToChat(conn, text): # 检查输入是否是JSON格式(包含说话人信息) speaker_name = None @@ -102,12 +114,13 @@ async def no_voice_close_connect(conn, have_voice): async def max_out_size(conn): + # 播放超出最大输出字数的提示 + conn.client_abort = False text = "不好意思,我现在有点事情要忙,明天这个时候我们再聊,约好了哦!明天不见不散,拜拜!" await send_stt_message(conn, text) file_path = "config/assets/max_output_size.wav" - conn.tts.tts_audio_queue.put((SentenceType.FIRST, [], text)) - play_audio_frames(conn, file_path) - conn.tts.tts_audio_queue.put((SentenceType.LAST, [], None)) + opus_packets = audio_to_data(file_path) + conn.tts.tts_audio_queue.put((SentenceType.LAST, opus_packets, text)) conn.close_after_chat = True @@ -125,35 +138,25 @@ async def check_bind_device(conn): # 播放提示音 music_path = "config/assets/bind_code.wav" - conn.tts.tts_audio_queue.put((SentenceType.FIRST, [], text)) - play_audio_frames(conn, music_path) + opus_packets = audio_to_data(music_path) + conn.tts.tts_audio_queue.put((SentenceType.FIRST, opus_packets, text)) # 逐个播放数字 for i in range(6): # 确保只播放6位数字 try: digit = conn.bind_code[i] num_path = f"config/assets/bind_code/{digit}.wav" - play_audio_frames(conn, num_path) + num_packets = 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}") continue conn.tts.tts_audio_queue.put((SentenceType.LAST, [], None)) else: + # 播放未绑定提示 + conn.client_abort = False text = f"没有找到该设备的版本信息,请正确配置 OTA地址,然后重新编译固件。" await send_stt_message(conn, text) music_path = "config/assets/bind_not_found.wav" - conn.tts.tts_audio_queue.put((SentenceType.FIRST, [], text)) - play_audio_frames(conn, music_path) - conn.tts.tts_audio_queue.put((SentenceType.LAST, [], None)) - - -def play_audio_frames(conn, file_path): - """播放音频文件并处理发送帧数据""" - def handle_audio_frame(frame_data): - conn.tts.tts_audio_queue.put((SentenceType.MIDDLE, frame_data, None)) - - audio_to_data_stream( - file_path, - is_opus=True, - callback=handle_audio_frame - ) + opus_packets = audio_to_data(music_path) + conn.tts.tts_audio_queue.put((SentenceType.LAST, opus_packets, text)) diff --git a/main/xiaozhi-server/core/handle/sendAudioHandle.py b/main/xiaozhi-server/core/handle/sendAudioHandle.py index 8604279a..661c4e21 100644 --- a/main/xiaozhi-server/core/handle/sendAudioHandle.py +++ b/main/xiaozhi-server/core/handle/sendAudioHandle.py @@ -1,8 +1,9 @@ import json -import asyncio import time -from core.providers.tts.dto.dto import SentenceType +import asyncio from core.utils import textUtils +from core.utils.util import audio_to_data +from core.providers.tts.dto.dto import SentenceType TAG = __name__ @@ -30,32 +31,24 @@ async def sendAudioMessage(conn, sentenceType, audios, text): # 播放音频 -async def sendAudio(conn, audios, pre_buffer=False): +async def sendAudio(conn, audios, frame_duration=60): """ 发送单个opus包,支持流控 Args: conn: 连接对象 opus_packet: 单个opus数据包 pre_buffer: 快速发送音频 + frame_duration: 帧时长(毫秒),匹配 Opus 编码 """ - if audios is None: + if audios is None or len(audios) == 0: return if isinstance(audios, bytes): if conn.client_abort: return - # 短音频直接发送(例如:提示音) - if pre_buffer: - await conn.websocket.send(audios) - return - - # 重置没有声音的状态 conn.last_activity_time = time.time() * 1000 - # 流控逻辑:确保按60ms的帧时长间隔发送 - frame_duration = 60 # 毫秒 - # 获取或初始化流控状态 if not hasattr(conn, "audio_flow_control"): conn.audio_flow_control = { @@ -66,13 +59,10 @@ async def sendAudio(conn, audios, pre_buffer=False): flow_control = conn.audio_flow_control current_time = time.perf_counter() - - # 计算期望的发送时间 + # 计算预期发送时间 expected_time = flow_control["start_time"] + ( flow_control["packet_count"] * frame_duration / 1000 ) - - # 流控延迟 delay = expected_time - current_time if delay > 0: await asyncio.sleep(delay) @@ -83,6 +73,35 @@ async def sendAudio(conn, audios, pre_buffer=False): # 更新流控状态 flow_control["packet_count"] += 1 flow_control["last_send_time"] = time.perf_counter() + else: + # 文件型音频走普通播放 + start_time = time.perf_counter() + play_position = 0 + + # 执行预缓冲 + pre_buffer_frames = min(3, len(audios)) + for i in range(pre_buffer_frames): + await conn.websocket.send(audios[i]) + remaining_audios = audios[pre_buffer_frames:] + + # 播放剩余音频帧 + for opus_packet in remaining_audios: + if conn.client_abort: + break + + # 重置没有声音的状态 + conn.last_activity_time = time.time() * 1000 + + # 计算预期发送时间 + expected_time = start_time + (play_position / 1000) + current_time = time.perf_counter() + delay = expected_time - current_time + if delay > 0: + await asyncio.sleep(delay) + + await conn.websocket.send(opus_packet) + + play_position += frame_duration async def send_tts_message(conn, state, text=None): @@ -101,12 +120,8 @@ 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" ) - conn.tts.audio_to_opus_data_stream( - stop_tts_notify_voice, - callback=lambda audio_data: asyncio.run_coroutine_threadsafe( - sendAudio(conn, audio_data, True), conn.loop - ), - ) + audios = audio_to_data(stop_tts_notify_voice, is_opus=True) + await sendAudio(conn, audios) # 清除服务端讲话状态 conn.clearSpeakStatus() diff --git a/main/xiaozhi-server/core/handle/textHandle.py b/main/xiaozhi-server/core/handle/textHandle.py index 86363217..b5e87783 100644 --- a/main/xiaozhi-server/core/handle/textHandle.py +++ b/main/xiaozhi-server/core/handle/textHandle.py @@ -1,154 +1,14 @@ -import json -import time -from core.handle.abortHandle import handleAbortMessage -from core.handle.helloHandle import handleHelloMessage -from core.providers.tools.device_mcp import handle_mcp_message -from core.utils.util import remove_punctuation_and_length, filter_sensitive_info -from core.handle.receiveAudioHandle import startToChat, handleAudioMessage -from core.providers.tools.device_iot import handleIotDescriptors, handleIotStatus -from core.handle.reportHandle import enqueue_asr_report -import asyncio +from core.handle.textMessageHandlerRegistry import TextMessageHandlerRegistry +from core.handle.textMessageProcessor import TextMessageProcessor TAG = __name__ +# 全局处理器注册表 +message_registry = TextMessageHandlerRegistry() + +# 创建全局消息处理器实例 +message_processor = TextMessageProcessor(message_registry) async def handleTextMessage(conn, message): """处理文本消息""" - try: - msg_json = json.loads(message) - if isinstance(msg_json, int): - conn.logger.bind(tag=TAG).info(f"收到文本消息:{message}") - await conn.websocket.send(message) - return - if msg_json["type"] == "hello": - conn.logger.bind(tag=TAG).info(f"收到hello消息:{message}") - await handleHelloMessage(conn, msg_json) - elif msg_json["type"] == "abort": - conn.logger.bind(tag=TAG).info(f"收到abort消息:{message}") - await handleAbortMessage(conn) - elif msg_json["type"] == "listen": - conn.logger.bind(tag=TAG).info(f"收到listen消息:{message}") - if "mode" in msg_json: - conn.client_listen_mode = msg_json["mode"] - conn.logger.bind(tag=TAG).debug( - f"客户端拾音模式:{conn.client_listen_mode}" - ) - if msg_json["state"] == "start": - conn.client_have_voice = True - conn.client_voice_stop = False - elif msg_json["state"] == "stop": - conn.client_have_voice = True - conn.client_voice_stop = True - if len(conn.asr_audio) > 0: - await handleAudioMessage(conn, b"") - elif msg_json["state"] == "detect": - conn.client_have_voice = False - conn.asr_audio.clear() - if "text" in msg_json: - conn.last_activity_time = time.time() * 1000 - original_text = msg_json["text"] # 保留原始文本 - filtered_len, filtered_text = remove_punctuation_and_length( - original_text - ) - # 识别是否是唤醒词 - is_wakeup_words = filtered_text in conn.config.get("wakeup_words") - if not is_wakeup_words: - # 上报纯文字数据(复用ASR上报功能,但不提供音频数据) - enqueue_asr_report(conn, original_text, []) - # 否则需要LLM对文字内容进行答复 - await startToChat(conn, original_text) - elif msg_json["type"] == "iot": - conn.logger.bind(tag=TAG).info(f"收到iot消息:{message}") - if "descriptors" in msg_json: - asyncio.create_task(handleIotDescriptors(conn, msg_json["descriptors"])) - if "states" in msg_json: - asyncio.create_task(handleIotStatus(conn, msg_json["states"])) - elif msg_json["type"] == "mcp": - conn.logger.bind(tag=TAG).info(f"收到mcp消息:{message[:100]}") - if "payload" in msg_json: - asyncio.create_task( - handle_mcp_message(conn, conn.mcp_client, msg_json["payload"]) - ) - elif msg_json["type"] == "server": - # 记录日志时过滤敏感信息 - conn.logger.bind(tag=TAG).info( - f"收到服务器消息:{filter_sensitive_info(msg_json)}" - ) - # 如果配置是从API读取的,则需要验证secret - if not conn.read_config_from_api: - return - # 获取post请求的secret - post_secret = msg_json.get("content", {}).get("secret", "") - secret = conn.config["manager-api"].get("secret", "") - # 如果secret不匹配,则返回 - if post_secret != secret: - await conn.websocket.send( - json.dumps( - { - "type": "server", - "status": "error", - "message": "服务器密钥验证失败", - } - ) - ) - return - # 动态更新配置 - if msg_json["action"] == "update_config": - try: - # 更新WebSocketServer的配置 - if not conn.server: - await conn.websocket.send( - json.dumps( - { - "type": "server", - "status": "error", - "message": "无法获取服务器实例", - "content": {"action": "update_config"}, - } - ) - ) - return - - if not await conn.server.update_config(): - await conn.websocket.send( - json.dumps( - { - "type": "server", - "status": "error", - "message": "更新服务器配置失败", - "content": {"action": "update_config"}, - } - ) - ) - return - - # 发送成功响应 - await conn.websocket.send( - json.dumps( - { - "type": "server", - "status": "success", - "message": "配置更新成功", - "content": {"action": "update_config"}, - } - ) - ) - except Exception as e: - conn.logger.bind(tag=TAG).error(f"更新配置失败: {str(e)}") - await conn.websocket.send( - json.dumps( - { - "type": "server", - "status": "error", - "message": f"更新配置失败: {str(e)}", - "content": {"action": "update_config"}, - } - ) - ) - # 重启服务器 - elif msg_json["action"] == "restart": - await conn.handle_restart(msg_json) - else: - conn.logger.bind(tag=TAG).error(f"收到未知类型消息:{message}") - except json.JSONDecodeError: - await conn.websocket.send(message) + await message_processor.process_message(conn, message) diff --git a/main/xiaozhi-server/core/handle/textHandler/abortMessageHandler.py b/main/xiaozhi-server/core/handle/textHandler/abortMessageHandler.py new file mode 100644 index 00000000..dc540d24 --- /dev/null +++ b/main/xiaozhi-server/core/handle/textHandler/abortMessageHandler.py @@ -0,0 +1,16 @@ +from typing import Dict, Any + +from core.handle.abortHandle import handleAbortMessage +from core.handle.textMessageHandler import TextMessageHandler +from core.handle.textMessageType import TextMessageType + + +class AbortTextMessageHandler(TextMessageHandler): + """Abort消息处理器""" + + @property + def message_type(self) -> TextMessageType: + return TextMessageType.ABORT + + async def handle(self, conn, msg_json: Dict[str, Any]) -> None: + await handleAbortMessage(conn) diff --git a/main/xiaozhi-server/core/handle/textHandler/helloMessageHandler.py b/main/xiaozhi-server/core/handle/textHandler/helloMessageHandler.py new file mode 100644 index 00000000..1839814e --- /dev/null +++ b/main/xiaozhi-server/core/handle/textHandler/helloMessageHandler.py @@ -0,0 +1,16 @@ +from typing import Dict, Any + +from core.handle.helloHandle import handleHelloMessage +from core.handle.textMessageHandler import TextMessageHandler +from core.handle.textMessageType import TextMessageType + + +class HelloTextMessageHandler(TextMessageHandler): + """Hello消息处理器""" + + @property + def message_type(self) -> TextMessageType: + return TextMessageType.HELLO + + async def handle(self, conn, msg_json: Dict[str, Any]) -> None: + await handleHelloMessage(conn, msg_json) \ No newline at end of file diff --git a/main/xiaozhi-server/core/handle/textHandler/iotMessageHandler.py b/main/xiaozhi-server/core/handle/textHandler/iotMessageHandler.py new file mode 100644 index 00000000..335d08b0 --- /dev/null +++ b/main/xiaozhi-server/core/handle/textHandler/iotMessageHandler.py @@ -0,0 +1,20 @@ +import asyncio +from typing import Dict, Any + +from core.handle.textMessageHandler import TextMessageHandler +from core.handle.textMessageType import TextMessageType +from core.providers.tools.device_iot import handleIotStatus, handleIotDescriptors + + +class IotTextMessageHandler(TextMessageHandler): + """IOT消息处理器""" + + @property + def message_type(self) -> TextMessageType: + return TextMessageType.IOT + + async def handle(self, conn, msg_json: Dict[str, Any]) -> None: + if "descriptors" in msg_json: + asyncio.create_task(handleIotDescriptors(conn, msg_json["descriptors"])) + if "states" in msg_json: + asyncio.create_task(handleIotStatus(conn, msg_json["states"])) \ No newline at end of file diff --git a/main/xiaozhi-server/core/handle/textHandler/listenMessageHandler.py b/main/xiaozhi-server/core/handle/textHandler/listenMessageHandler.py new file mode 100644 index 00000000..97286dfe --- /dev/null +++ b/main/xiaozhi-server/core/handle/textHandler/listenMessageHandler.py @@ -0,0 +1,63 @@ +import time +from typing import Dict, Any + +from core.handle.receiveAudioHandle import handleAudioMessage, 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 + +TAG = __name__ + +class ListenTextMessageHandler(TextMessageHandler): + """Listen消息处理器""" + + @property + def message_type(self) -> TextMessageType: + return TextMessageType.LISTEN + + async def handle(self, conn, msg_json: Dict[str, Any]) -> None: + if "mode" in msg_json: + conn.client_listen_mode = msg_json["mode"] + conn.logger.bind(tag=TAG).debug( + f"客户端拾音模式:{conn.client_listen_mode}" + ) + if msg_json["state"] == "start": + conn.client_have_voice = True + conn.client_voice_stop = False + elif msg_json["state"] == "stop": + conn.client_have_voice = True + conn.client_voice_stop = True + if len(conn.asr_audio) > 0: + await handleAudioMessage(conn, b"") + elif msg_json["state"] == "detect": + conn.client_have_voice = False + conn.asr_audio.clear() + if "text" in msg_json: + conn.last_activity_time = time.time() * 1000 + original_text = msg_json["text"] # 保留原始文本 + filtered_len, filtered_text = remove_punctuation_and_length( + original_text + ) + + # 识别是否是唤醒词 + is_wakeup_words = filtered_text in conn.config.get("wakeup_words") + # 是否开启唤醒词回复 + enable_greeting = conn.config.get("enable_greeting", True) + + if is_wakeup_words and not enable_greeting: + # 如果是唤醒词,且关闭了唤醒词回复,就不用回答 + await send_stt_message(conn, original_text) + await send_tts_message(conn, "stop", None) + conn.client_is_speaking = False + elif is_wakeup_words: + conn.just_woken_up = True + # 上报纯文字数据(复用ASR上报功能,但不提供音频数据) + enqueue_asr_report(conn, "嘿,你好呀", []) + await startToChat(conn, "嘿,你好呀") + else: + # 上报纯文字数据(复用ASR上报功能,但不提供音频数据) + enqueue_asr_report(conn, original_text, []) + # 否则需要LLM对文字内容进行答复 + await startToChat(conn, original_text) \ No newline at end of file diff --git a/main/xiaozhi-server/core/handle/textHandler/mcpMessageHandler.py b/main/xiaozhi-server/core/handle/textHandler/mcpMessageHandler.py new file mode 100644 index 00000000..65876f24 --- /dev/null +++ b/main/xiaozhi-server/core/handle/textHandler/mcpMessageHandler.py @@ -0,0 +1,20 @@ +import asyncio +from typing import Dict, Any + +from core.handle.textMessageHandler import TextMessageHandler +from core.handle.textMessageType import TextMessageType +from core.providers.tools.device_mcp import handle_mcp_message + + +class McpTextMessageHandler(TextMessageHandler): + """MCP消息处理器""" + + @property + def message_type(self) -> TextMessageType: + return TextMessageType.MCP + + async def handle(self, conn, msg_json: Dict[str, Any]) -> None: + if "payload" in msg_json: + asyncio.create_task( + handle_mcp_message(conn, conn.mcp_client, msg_json["payload"]) + ) \ No newline at end of file diff --git a/main/xiaozhi-server/core/handle/textHandler/serverMessageHandler.py b/main/xiaozhi-server/core/handle/textHandler/serverMessageHandler.py new file mode 100644 index 00000000..b9a23588 --- /dev/null +++ b/main/xiaozhi-server/core/handle/textHandler/serverMessageHandler.py @@ -0,0 +1,92 @@ +import asyncio +import json +from typing import Dict, Any + +from core.handle.textMessageHandler import TextMessageHandler +from core.handle.textMessageType import TextMessageType +from core.providers.tools.device_mcp import handle_mcp_message + +TAG = __name__ + +class ServerTextMessageHandler(TextMessageHandler): + """MCP消息处理器""" + + @property + def message_type(self) -> TextMessageType: + return TextMessageType.SERVER + + async def handle(self, conn, msg_json: Dict[str, Any]) -> None: + # 如果配置是从API读取的,则需要验证secret + if not conn.read_config_from_api: + return + # 获取post请求的secret + post_secret = msg_json.get("content", {}).get("secret", "") + secret = conn.config["manager-api"].get("secret", "") + # 如果secret不匹配,则返回 + if post_secret != secret: + await conn.websocket.send( + json.dumps( + { + "type": "server", + "status": "error", + "message": "服务器密钥验证失败", + } + ) + ) + return + # 动态更新配置 + if msg_json["action"] == "update_config": + try: + # 更新WebSocketServer的配置 + if not conn.server: + await conn.websocket.send( + json.dumps( + { + "type": "server", + "status": "error", + "message": "无法获取服务器实例", + "content": {"action": "update_config"}, + } + ) + ) + return + + if not await conn.server.update_config(): + await conn.websocket.send( + json.dumps( + { + "type": "server", + "status": "error", + "message": "更新服务器配置失败", + "content": {"action": "update_config"}, + } + ) + ) + return + + # 发送成功响应 + await conn.websocket.send( + json.dumps( + { + "type": "server", + "status": "success", + "message": "配置更新成功", + "content": {"action": "update_config"}, + } + ) + ) + except Exception as e: + conn.logger.bind(tag=TAG).error(f"更新配置失败: {str(e)}") + await conn.websocket.send( + json.dumps( + { + "type": "server", + "status": "error", + "message": f"更新配置失败: {str(e)}", + "content": {"action": "update_config"}, + } + ) + ) + # 重启服务器 + elif msg_json["action"] == "restart": + await conn.handle_restart(msg_json) \ No newline at end of file diff --git a/main/xiaozhi-server/core/handle/textMessageHandler.py b/main/xiaozhi-server/core/handle/textMessageHandler.py new file mode 100644 index 00000000..f94a0bac --- /dev/null +++ b/main/xiaozhi-server/core/handle/textMessageHandler.py @@ -0,0 +1,21 @@ +from abc import abstractmethod, ABC +from typing import Dict, Any + +from core.handle.textMessageType import TextMessageType + +TAG = __name__ + + +class TextMessageHandler(ABC): + """消息处理器抽象基类""" + + @abstractmethod + async def handle(self, conn, msg_json: Dict[str, Any]) -> None: + """处理消息的抽象方法""" + pass + + @property + @abstractmethod + def message_type(self) -> TextMessageType: + """返回处理的消息类型""" + pass diff --git a/main/xiaozhi-server/core/handle/textMessageHandlerRegistry.py b/main/xiaozhi-server/core/handle/textMessageHandlerRegistry.py new file mode 100644 index 00000000..e90d7231 --- /dev/null +++ b/main/xiaozhi-server/core/handle/textMessageHandlerRegistry.py @@ -0,0 +1,45 @@ +from typing import Dict, Optional + +from core.handle.textHandler.abortMessageHandler import AbortTextMessageHandler +from core.handle.textHandler.helloMessageHandler import HelloTextMessageHandler +from core.handle.textHandler.iotMessageHandler import IotTextMessageHandler +from core.handle.textHandler.listenMessageHandler import ListenTextMessageHandler +from core.handle.textHandler.mcpMessageHandler import McpTextMessageHandler +from core.handle.textMessageHandler import TextMessageHandler +from core.handle.textHandler.serverMessageHandler import ServerTextMessageHandler + +TAG = __name__ + + +class TextMessageHandlerRegistry: + """消息处理器注册表""" + + def __init__(self): + self._handlers: Dict[str, TextMessageHandler] = {} + self._register_default_handlers() + + def _register_default_handlers(self) -> None: + """注册默认的消息处理器""" + handlers = [ + HelloTextMessageHandler(), + AbortTextMessageHandler(), + ListenTextMessageHandler(), + IotTextMessageHandler(), + McpTextMessageHandler(), + ServerTextMessageHandler(), + ] + + for handler in handlers: + self.register_handler(handler) + + def register_handler(self, handler: TextMessageHandler) -> None: + """注册消息处理器""" + self._handlers[handler.message_type.value] = handler + + def get_handler(self, message_type: str) -> Optional[TextMessageHandler]: + """获取消息处理器""" + return self._handlers.get(message_type) + + def get_supported_types(self) -> list: + """获取支持的消息类型""" + return list(self._handlers.keys()) diff --git a/main/xiaozhi-server/core/handle/textMessageProcessor.py b/main/xiaozhi-server/core/handle/textMessageProcessor.py new file mode 100644 index 00000000..0cae5e09 --- /dev/null +++ b/main/xiaozhi-server/core/handle/textMessageProcessor.py @@ -0,0 +1,41 @@ +import json + +from core.handle.textMessageHandlerRegistry import TextMessageHandlerRegistry + +TAG = __name__ + + +class TextMessageProcessor: + """消息处理器主类""" + + def __init__(self, registry: TextMessageHandlerRegistry): + self.registry = registry + + async def process_message(self, conn, message: str) -> None: + """处理消息的主入口""" + try: + # 解析JSON消息 + msg_json = json.loads(message) + + # 处理JSON消息 + if isinstance(msg_json, dict): + message_type = msg_json.get("type") + + # 记录日志 + conn.logger.bind(tag=TAG).info(f"收到{message_type}消息:{message}") + + # 获取并执行处理器 + handler = self.registry.get_handler(message_type) + if handler: + await handler.handle(conn, msg_json) + else: + conn.logger.bind(tag=TAG).error(f"收到未知类型消息:{message}") + # 处理纯数字消息 + elif isinstance(msg_json, int): + conn.logger.bind(tag=TAG).info(f"收到数字消息:{message}") + await conn.websocket.send(message) + + except json.JSONDecodeError: + # 非JSON消息直接转发 + conn.logger.bind(tag=TAG).error(f"解析到错误的消息:{message}") + await conn.websocket.send(message) diff --git a/main/xiaozhi-server/core/handle/textMessageType.py b/main/xiaozhi-server/core/handle/textMessageType.py new file mode 100644 index 00000000..53e71b71 --- /dev/null +++ b/main/xiaozhi-server/core/handle/textMessageType.py @@ -0,0 +1,11 @@ +from enum import Enum + + +class TextMessageType(Enum): + """消息类型枚举""" + HELLO = "hello" + ABORT = "abort" + LISTEN = "listen" + IOT = "iot" + MCP = "mcp" + SERVER = "server" diff --git a/main/xiaozhi-server/core/providers/asr/base.py b/main/xiaozhi-server/core/providers/asr/base.py index f4803d31..250a25f2 100644 --- a/main/xiaozhi-server/core/providers/asr/base.py +++ b/main/xiaozhi-server/core/providers/asr/base.py @@ -1,14 +1,14 @@ import os +import io import wave import uuid +import json +import time import queue import asyncio import traceback import threading import opuslib_next -import json -import io -import time import concurrent.futures from abc import ABC, abstractmethod from config.logger import setup_logging @@ -87,11 +87,9 @@ class ASRProviderBase(ABC): # 预先准备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() @@ -149,7 +147,7 @@ class ASRProviderBase(ABC): # 处理结果 - raw_text, file_path = results.get("asr", ("", None)) + raw_text, _ = results.get("asr", ("", None)) speaker_name = results.get("voiceprint", None) # 记录识别结果 diff --git a/main/xiaozhi-server/core/providers/llm/AliBL/AliBL.py b/main/xiaozhi-server/core/providers/llm/AliBL/AliBL.py index 012035ee..a68c4348 100644 --- a/main/xiaozhi-server/core/providers/llm/AliBL/AliBL.py +++ b/main/xiaozhi-server/core/providers/llm/AliBL/AliBL.py @@ -1,8 +1,10 @@ from config.logger import setup_logging from http import HTTPStatus +import dashscope from dashscope import Application from core.providers.llm.base import LLMProviderBase from core.utils.util import check_model_key +import time TAG = __name__ logger = setup_logging() @@ -15,6 +17,7 @@ class LLMProvider(LLMProviderBase): self.base_url = config.get("base_url") self.is_No_prompt = config.get("is_no_prompt") self.memory_id = config.get("ali_memory_id") + self.streaming_chunk_size = config.get("streaming_chunk_size", 3) # 每次流式返回的字符数 check_model_key("AliBLLLM", self.api_key) def response(self, session_id, dialogue): @@ -32,6 +35,8 @@ class LLMProvider(LLMProviderBase): "app_id": self.app_id, "session_id": session_id, "messages": dialogue, + # 开启SDK原生流式 + "stream": True, } if self.memory_id != False: # 百练memory需要prompt参数 @@ -42,25 +47,63 @@ class LLMProvider(LLMProviderBase): f"【阿里百练API服务】处理后的prompt: {prompt}" ) + # 可选地设置自定义API基地址(若配置为兼容模式URL则忽略) + if self.base_url and ("/api/" in self.base_url): + dashscope.base_http_api_url = self.base_url + responses = Application.call(**call_params) - if responses.status_code != HTTPStatus.OK: - logger.bind(tag=TAG).error( - f"code={responses.status_code}, " - f"message={responses.message}, " - f"请参考文档:https://help.aliyun.com/zh/model-studio/developer-reference/error-code" - ) - yield "【阿里百练API服务响应异常】" - else: - logger.bind(tag=TAG).debug( - f"【阿里百练API服务】构造参数: {call_params}" - ) - yield responses.output.text + + # 流式处理(SDK在stream=True时返回可迭代对象;否则返回单次响应对象) + logger.bind(tag=TAG).debug( + f"【阿里百练API服务】构造参数: {dict(call_params, api_key='***')}" + ) + + last_text = "" + try: + for resp in responses: + if resp.status_code != HTTPStatus.OK: + logger.bind(tag=TAG).error( + f"code={resp.status_code}, message={resp.message}, 请参考文档:https://help.aliyun.com/zh/model-studio/developer-reference/error-code" + ) + continue + current_text = getattr(getattr(resp, "output", None), "text", None) + if current_text is None: + continue + # SDK流式为增量覆盖,计算差量输出 + if len(current_text) >= len(last_text): + delta = current_text[len(last_text):] + else: + # 避免偶发回退 + delta = current_text + if delta: + yield delta + last_text = current_text + except TypeError: + # 非流式回落(一次性返回) + if responses.status_code != HTTPStatus.OK: + logger.bind(tag=TAG).error( + f"code={responses.status_code}, message={responses.message}, 请参考文档:https://help.aliyun.com/zh/model-studio/developer-reference/error-code" + ) + yield "【阿里百练API服务响应异常】" + else: + full_text = getattr(getattr(responses, "output", None), "text", "") + logger.bind(tag=TAG).info( + f"【阿里百练API服务】完整响应长度: {len(full_text)}" + ) + for i in range(0, len(full_text), self.streaming_chunk_size): + chunk = full_text[i:i + self.streaming_chunk_size] + if chunk: + yield chunk except Exception as e: logger.bind(tag=TAG).error(f"【阿里百练API服务】响应异常: {e}") yield "【LLM服务响应异常】" def response_with_functions(self, session_id, dialogue, functions=None): - logger.bind(tag=TAG).error( - f"阿里百练暂未实现完整的工具调用(function call),建议使用其他意图识别" + # 阿里百练当前未支持原生的 function call。为保持兼容,这里回退到普通文本流式输出。 + # 上层会按 (content, tool_calls) 的形式消费,这里始终返回 (token, None) + logger.bind(tag=TAG).warning( + "阿里百练未实现原生 function call,已回退为纯文本流式输出" ) + for token in self.response(session_id, dialogue): + yield token, None diff --git a/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_executor.py b/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_executor.py index 9ae15d4b..ee0bcb2a 100644 --- a/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_executor.py +++ b/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_executor.py @@ -18,8 +18,8 @@ class ServerMCPExecutor(ToolExecutor): """初始化MCP管理器""" if not self._initialized: self.mcp_manager = ServerMCPManager(self.conn) - await self.mcp_manager.initialize_servers() self._initialized = True + await self.mcp_manager.initialize_servers() async def execute( self, conn, tool_name: str, arguments: Dict[str, Any] 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 52ab2b79..cf3deb21 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 @@ -68,6 +68,9 @@ class ServerMCPManager: # 输出当前支持的服务端MCP工具列表 if hasattr(self.conn, "func_handler") and self.conn.func_handler: + # 刷新工具缓存以确保服务端MCP工具被正确加载 + if hasattr(self.conn.func_handler, "tool_manager"): + self.conn.func_handler.tool_manager.refresh_tools() self.conn.func_handler.current_support_functions() def get_all_tools(self) -> List[Dict[str, Any]]: diff --git a/main/xiaozhi-server/core/providers/tts/aliyun_stream.py b/main/xiaozhi-server/core/providers/tts/aliyun_stream.py index 734681e2..38e70927 100644 --- a/main/xiaozhi-server/core/providers/tts/aliyun_stream.py +++ b/main/xiaozhi-server/core/providers/tts/aliyun_stream.py @@ -478,3 +478,142 @@ class TTSProvider(TTSProviderBase): finally: self._monitor_task = None + def to_tts(self, text: str) -> list: + """非流式TTS处理,用于测试及保存音频文件的场景""" + try: + # 创建新的事件循环 + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + + # 生成会话ID + session_id = uuid.uuid4().hex + # 存储音频数据 + audio_data = [] + + async def _generate_audio(): + # 刷新Token(如果需要) + if self._is_token_expired(): + self._refresh_token() + + # 建立WebSocket连接 + ws = await websockets.connect( + self.ws_url, + additional_headers={"X-NLS-Token": self.token}, + ping_interval=30, + ping_timeout=10, + close_timeout=10, + ) + try: + # 发送StartSynthesis请求 + start_message_id = str(uuid.uuid4().hex) + start_request = { + "header": { + "message_id": start_message_id, + "task_id": session_id, + "namespace": "FlowingSpeechSynthesizer", + "name": "StartSynthesis", + "appkey": self.appkey, + }, + "payload": { + "voice": self.voice, + "format": self.format, + "sample_rate": self.sample_rate, + "volume": self.volume, + "speech_rate": self.speech_rate, + "pitch_rate": self.pitch_rate, + "enable_subtitle": True, + }, + } + await ws.send(json.dumps(start_request)) + + # 等待SynthesisStarted响应 + synthesis_started = False + while not synthesis_started: + msg = await ws.recv() + if isinstance(msg, str): + data = json.loads(msg) + header = data.get("header", {}) + if header.get("name") == "SynthesisStarted": + synthesis_started = True + logger.bind(tag=TAG).debug("TTS合成已启动") + elif header.get("name") == "TaskFailed": + error_info = data.get("payload", {}).get( + "error_info", {} + ) + error_code = error_info.get("error_code") + error_message = error_info.get( + "error_message", "未知错误" + ) + raise Exception( + f"启动合成失败: {error_code} - {error_message}" + ) + + # 发送文本合成请求 + filtered_text = MarkdownCleaner.clean_markdown(text) + run_message_id = str(uuid.uuid4().hex) + run_request = { + "header": { + "message_id": run_message_id, + "task_id": session_id, + "namespace": "FlowingSpeechSynthesizer", + "name": "RunSynthesis", + "appkey": self.appkey, + }, + "payload": {"text": filtered_text}, + } + await ws.send(json.dumps(run_request)) + + # 发送停止合成请求 + stop_message_id = str(uuid.uuid4().hex) + stop_request = { + "header": { + "message_id": stop_message_id, + "task_id": session_id, + "namespace": "FlowingSpeechSynthesizer", + "name": "StopSynthesis", + "appkey": self.appkey, + } + } + await ws.send(json.dumps(stop_request)) + + # 接收音频数据 + synthesis_completed = False + while not synthesis_completed: + msg = await ws.recv() + if isinstance(msg, (bytes, bytearray)): + self.opus_encoder.encode_pcm_to_opus_stream( + msg, + end_of_stream=False, + callback=lambda opus: audio_data.append(opus) + ) + elif isinstance(msg, str): + data = json.loads(msg) + header = data.get("header", {}) + event_name = header.get("name") + if event_name == "SynthesisCompleted": + synthesis_completed = True + logger.bind(tag=TAG).debug("TTS合成完成") + elif event_name == "TaskFailed": + error_info = data.get("payload", {}).get( + "error_info", {} + ) + error_code = error_info.get("error_code") + error_message = error_info.get( + "error_message", "未知错误" + ) + raise Exception( + f"合成失败: {error_code} - {error_message}" + ) + finally: + try: + await ws.close() + except: + pass + + loop.run_until_complete(_generate_audio()) + loop.close() + + return audio_data + except Exception as e: + logger.bind(tag=TAG).error(f"生成音频数据失败: {str(e)}") + return [] \ No newline at end of file diff --git a/main/xiaozhi-server/core/providers/tts/base.py b/main/xiaozhi-server/core/providers/tts/base.py index 14301205..04a7fa36 100644 --- a/main/xiaozhi-server/core/providers/tts/base.py +++ b/main/xiaozhi-server/core/providers/tts/base.py @@ -1,21 +1,22 @@ import os import re -import queue +import time import uuid +import queue import asyncio import threading -from typing import Callable, Any +import traceback from core.utils import p3 -import time from datetime import datetime from core.utils import textUtils +from typing import Callable, Any from abc import ABC, abstractmethod from config.logger import setup_logging -from core.utils.util import audio_bytes_to_data_stream, audio_to_data_stream from core.utils.tts import MarkdownCleaner from core.utils.output_counter import add_device_output from core.handle.reportHandle import enqueue_tts_report from core.handle.sendAudioHandle import sendAudioMessage +from core.utils.util import audio_bytes_to_data_stream, audio_to_data_stream from core.providers.tts.dto.dto import ( TTSMessageDTO, SentenceType, @@ -23,8 +24,6 @@ from core.providers.tts.dto.dto import ( InterfaceType, ) -import traceback - TAG = __name__ logger = setup_logging() @@ -144,6 +143,68 @@ class TTSProviderBase(ABC): except Exception as e: logger.bind(tag=TAG).error(f"Failed to generate TTS file: {e}") return None + + def to_tts(self, text): + text = MarkdownCleaner.clean_markdown(text) + max_repeat_time = 5 + if self.delete_audio_file: + # 需要删除文件的直接转为音频数据 + while max_repeat_time > 0: + try: + audio_bytes = asyncio.run(self.text_to_speak(text, None)) + if audio_bytes: + audio_datas = [] + audio_bytes_to_data_stream( + audio_bytes, + file_type=self.audio_file_type, + is_opus=True, + callback=lambda data: audio_datas.append(data) + ) + return audio_datas + else: + max_repeat_time -= 1 + except Exception as e: + logger.bind(tag=TAG).warning( + f"语音生成失败{5 - max_repeat_time + 1}次: {text},错误: {e}" + ) + max_repeat_time -= 1 + if max_repeat_time > 0: + logger.bind(tag=TAG).info( + f"语音生成成功: {text},重试{5 - max_repeat_time}次" + ) + else: + logger.bind(tag=TAG).error( + f"语音生成失败: {text},请检查网络或服务是否正常" + ) + return None + else: + tmp_file = self.generate_filename() + try: + while not os.path.exists(tmp_file) and max_repeat_time > 0: + try: + asyncio.run(self.text_to_speak(text, tmp_file)) + except Exception as e: + logger.bind(tag=TAG).warning( + f"语音生成失败{5 - max_repeat_time + 1}次: {text},错误: {e}" + ) + # 未执行成功,删除文件 + if os.path.exists(tmp_file): + os.remove(tmp_file) + max_repeat_time -= 1 + + if max_repeat_time > 0: + logger.bind(tag=TAG).info( + f"语音生成成功: {text}:{tmp_file},重试{5 - max_repeat_time}次" + ) + else: + logger.bind(tag=TAG).error( + f"语音生成失败: {text},请检查网络或服务是否正常" + ) + + return tmp_file + except Exception as e: + logger.bind(tag=TAG).error(f"Failed to generate TTS file: {e}") + return None @abstractmethod async def text_to_speak(self, text, output_file): @@ -284,8 +345,8 @@ class TTSProviderBase(ABC): enqueue_audio = [] enqueue_text = text - # 计算音频数据的帧数 - if isinstance(audio_datas, bytes): + # 收集上报音频数据 + if isinstance(audio_datas, bytes) and enqueue_audio is not None: enqueue_audio.append(audio_datas) # 发送音频 diff --git a/main/xiaozhi-server/core/providers/tts/fishspeech.py b/main/xiaozhi-server/core/providers/tts/fishspeech.py index bbf19164..83b3bc3f 100644 --- a/main/xiaozhi-server/core/providers/tts/fishspeech.py +++ b/main/xiaozhi-server/core/providers/tts/fishspeech.py @@ -143,8 +143,8 @@ class TTSProvider(TTSProviderBase): data = { "text": text, "references": [ - ServeReferenceAudio(audio=audio if audio else b"", text=text) - for text, audio in zip(ref_texts, byte_audios) + ServeReferenceAudio(audio=audio if audio else b"", text=ref_text) + for ref_text, audio in zip(ref_texts, byte_audios) ], "reference_id": self.reference_id, "normalize": self.normalize, diff --git a/main/xiaozhi-server/core/providers/tts/huoshan_double_stream.py b/main/xiaozhi-server/core/providers/tts/huoshan_double_stream.py index 463c16fc..ff295581 100644 --- a/main/xiaozhi-server/core/providers/tts/huoshan_double_stream.py +++ b/main/xiaozhi-server/core/providers/tts/huoshan_double_stream.py @@ -628,3 +628,104 @@ class TTSProvider(TTSProviderBase): def wav_to_opus_data_audio_raw_stream(self, raw_data_var, is_end=False, callback: Callable[[Any], Any]=None): return self.opus_encoder.encode_pcm_to_opus_stream(raw_data_var, is_end, callback=callback) + + def to_tts(self, text: str) -> list: + """非流式生成音频数据,用于生成音频及测试场景 + Args: + text: 要转换的文本 + Returns: + list: 音频数据列表 + """ + try: + # 创建事件循环 + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + + # 生成会话ID + session_id = uuid.uuid4().__str__().replace("-", "") + + # 存储音频数据 + audio_data = [] + + async def _generate_audio(): + # 创建新的WebSocket连接 + ws_header = { + "X-Api-App-Key": self.appId, + "X-Api-Access-Key": self.access_token, + "X-Api-Resource-Id": self.resource_id, + "X-Api-Connect-Id": uuid.uuid4(), + } + ws = await websockets.connect( + self.ws_url, additional_headers=ws_header, max_size=1000000000 + ) + + try: + # 启动会话 + header = Header( + message_type=FULL_CLIENT_REQUEST, + message_type_specific_flags=MsgTypeFlagWithEvent, + serial_method=JSON, + ).as_bytes() + optional = Optional( + event=EVENT_StartSession, sessionId=session_id + ).as_bytes() + payload = self.get_payload_bytes( + event=EVENT_StartSession, speaker=self.voice + ) + await self.send_event(ws, header, optional, payload) + + # 发送文本 + header = Header( + message_type=FULL_CLIENT_REQUEST, + message_type_specific_flags=MsgTypeFlagWithEvent, + serial_method=JSON, + ).as_bytes() + optional = Optional( + event=EVENT_TaskRequest, sessionId=session_id + ).as_bytes() + payload = self.get_payload_bytes( + event=EVENT_TaskRequest, text=text, speaker=self.voice + ) + await self.send_event(ws, header, optional, payload) + + # 发送结束会话请求 + header = Header( + message_type=FULL_CLIENT_REQUEST, + message_type_specific_flags=MsgTypeFlagWithEvent, + serial_method=JSON, + ).as_bytes() + optional = Optional( + event=EVENT_FinishSession, sessionId=session_id + ).as_bytes() + payload = str.encode("{}") + await self.send_event(ws, header, optional, payload) + + # 接收音频数据 + while True: + msg = await ws.recv() + res = self.parser_response(msg) + + if ( + res.optional.event == EVENT_TTSResponse + and res.header.message_type == AUDIO_ONLY_RESPONSE + ): + self.wav_to_opus_data_audio_raw_stream(res.payload, callback=lambda opus_frame: audio_data.append(opus_frame)) + elif res.optional.event == EVENT_SessionFinished: + break + + finally: + # 清理资源 + try: + await ws.close() + except: + pass + + # 运行异步任务 + loop.run_until_complete(_generate_audio()) + loop.close() + + return audio_data + + except Exception as e: + logger.bind(tag=TAG).error(f"生成音频数据失败: {str(e)}") + return [] diff --git a/main/xiaozhi-server/core/providers/tts/index_stream.py b/main/xiaozhi-server/core/providers/tts/index_stream.py index 6f6b829c..8abb06f2 100644 --- a/main/xiaozhi-server/core/providers/tts/index_stream.py +++ b/main/xiaozhi-server/core/providers/tts/index_stream.py @@ -1,8 +1,10 @@ import os +import time import queue -import asyncio -import traceback import aiohttp +import asyncio +import requests +import traceback from config.logger import setup_logging from core.utils.tts import MarkdownCleaner from core.providers.tts.base import TTSProviderBase @@ -177,3 +179,57 @@ class TTSProvider(TTSProviderBase): await super().close() if hasattr(self, "opus_encoder"): self.opus_encoder.close() + + def to_tts(self, text: str) -> list: + """非流式TTS处理,用于测试及保存音频文件的场景 + Args: + text: 要转换的文本 + Returns: + list: 返回opus编码后的音频数据列表 + """ + start_time = time.time() + text = MarkdownCleaner.clean_markdown(text) + + payload = {"text": text, "character": self.character} + + try: + with requests.post(self.api_url, json=payload, timeout=5) as response: + if response.status_code != 200: + logger.bind(tag=TAG).error( + f"TTS请求失败: {response.status_code}, {response.text}" + ) + return [] + + logger.info(f"TTS请求成功: {text}, 耗时: {time.time() - start_time}秒") + + # 使用opus编码器处理PCM数据 + opus_datas = [] + pcm_data = response.content + + # 计算每帧的字节数 + frame_bytes = int( + self.opus_encoder.sample_rate + * self.opus_encoder.channels + * self.opus_encoder.frame_size_ms + / 1000 + * 2 + ) + + # 分帧处理PCM数据 + for i in range(0, len(pcm_data), frame_bytes): + frame = pcm_data[i : i + frame_bytes] + if len(frame) < frame_bytes: + # 最后一帧可能不足,用0填充 + frame = frame + b"\x00" * (frame_bytes - len(frame)) + + self.opus_encoder.encode_pcm_to_opus_stream( + frame, + end_of_stream=(i + frame_bytes >= len(pcm_data)), + callback=lambda opus: opus_datas.append(opus) + ) + + return opus_datas + + except Exception as e: + logger.bind(tag=TAG).error(f"TTS请求异常: {e}") + return [] \ No newline at end of file diff --git a/main/xiaozhi-server/core/providers/tts/linkerai.py b/main/xiaozhi-server/core/providers/tts/linkerai.py index 875ed0b1..9b5e946f 100644 --- a/main/xiaozhi-server/core/providers/tts/linkerai.py +++ b/main/xiaozhi-server/core/providers/tts/linkerai.py @@ -1,8 +1,10 @@ import os +import time import queue -import asyncio -import traceback import aiohttp +import asyncio +import requests +import traceback from config.logger import setup_logging from core.utils.tts import MarkdownCleaner from core.providers.tts.base import TTSProviderBase @@ -109,10 +111,6 @@ class TTSProvider(TTSProviderBase): finally: return None - ################################################################################### - # linkerai单流式TTS重写父类的方法--结束 - ################################################################################### - async def text_to_speak(self, text, is_last): """流式处理TTS音频,每句只推送一次音频列表""" await self._tts_request(text, is_last) @@ -199,3 +197,71 @@ class TTSProvider(TTSProviderBase): except Exception as e: logger.bind(tag=TAG).error(f"TTS请求异常: {e}") self.tts_audio_queue.put((SentenceType.LAST, [], None)) + + def to_tts(self, text: str) -> list: + """非流式TTS处理,用于测试及保存音频文件的场景 + Args: + text: 要转换的文本 + Returns: + list: 返回opus编码后的音频数据列表 + """ + start_time = time.time() + text = MarkdownCleaner.clean_markdown(text) + + params = { + "tts_text": text, + "spk_id": self.voice, + "frame_duration": 60, + "stream": False, + "target_sr": 16000, + "audio_format": self.audio_format, + "instruct_text": "请生成一段自然流畅的语音", + } + headers = { + "Authorization": f"Bearer {self.access_token}", + "Content-Type": "application/json", + } + + try: + with requests.get( + self.api_url, params=params, headers=headers, timeout=5 + ) as response: + if response.status_code != 200: + logger.bind(tag=TAG).error( + f"TTS请求失败: {response.status_code}, {response.text}" + ) + return [] + + logger.info(f"TTS请求成功: {text}, 耗时: {time.time() - start_time}秒") + + # 使用opus编码器处理PCM数据 + opus_datas = [] + pcm_data = response.content + + # 计算每帧的字节数 + frame_bytes = int( + self.opus_encoder.sample_rate + * self.opus_encoder.channels + * self.opus_encoder.frame_size_ms + / 1000 + * 2 + ) + + # 分帧处理PCM数据 + for i in range(0, len(pcm_data), frame_bytes): + frame = pcm_data[i : i + frame_bytes] + if len(frame) < frame_bytes: + # 最后一帧可能不足,用0填充 + frame = frame + b"\x00" * (frame_bytes - len(frame)) + + self.opus_encoder.encode_pcm_to_opus_stream( + frame, + end_of_stream=(i + frame_bytes >= len(pcm_data)), + callback=lambda opus: opus_datas.append(opus) + ) + + return opus_datas + + except Exception as e: + logger.bind(tag=TAG).error(f"TTS请求异常: {e}") + return [] \ No newline at end of file diff --git a/main/xiaozhi-server/core/providers/tts/paddle_speech.py b/main/xiaozhi-server/core/providers/tts/paddle_speech.py index 7f1c6a23..864cf505 100644 --- a/main/xiaozhi-server/core/providers/tts/paddle_speech.py +++ b/main/xiaozhi-server/core/providers/tts/paddle_speech.py @@ -1,13 +1,15 @@ -import asyncio -import json -import base64 -import aiohttp -import numpy as np import io import wave +import json +import base64 +import asyncio import websockets -from core.providers.tts.base import TTSProviderBase +import numpy as np +from datetime import datetime from config.logger import setup_logging +from core.providers.tts.base import TTSProviderBase + + TAG = __name__ logger = setup_logging() @@ -18,11 +20,12 @@ class TTSProvider(TTSProviderBase): super().__init__(config, delete_audio_file) self.url = config.get("url", "ws://192.168.1.10:8092/paddlespeech/tts/streaming") self.protocol = config.get("protocol", "websocket") + if config.get("private_voice"): self.spk_id = int(config.get("private_voice")) else: - self.spk_id = int(config.get("spk_id", "0")) - + self.spk_id = int(config.get("spk_id", "0")) + sample_rate = config.get("sample_rate", 24000) self.sample_rate = float(sample_rate) if sample_rate else 24000 @@ -32,7 +35,21 @@ class TTSProvider(TTSProviderBase): volume = config.get("volume", 1.0) self.volume = float(volume) if volume else 1.0 - self.save_path = config.get("save_path", "./streaming_tts.wav") + self.delete_audio_file = config.get("delete_audio", True) + if not self.delete_audio_file: + timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") + save_path = config.get("save_path") + if save_path: + if not save_path.endswith('.wav'): + save_path = f"{save_path}_{timestamp}.wav" + else: + other_path = save_path[:-4] + save_path = f"{other_path}_{timestamp}.wav" + self.save_path = save_path + else: + self.save_path = f"./streaming_tts_{timestamp}.wav" + else: + self.save_path = None async def pcm_to_wav(self, pcm_data: bytes, sample_rate: int = 24000, num_channels: int = 1, bits_per_sample: int = 16) -> bytes: @@ -58,43 +75,9 @@ class TTSProvider(TTSProviderBase): async def text_to_speak(self, text, output_file): if self.protocol == "websocket": return await self.text_streaming(text, output_file) - elif self.protocol == "http": - return await self.text(text, output_file) else: raise ValueError("Unsupported protocol. Please use 'websocket' or 'http'.") - async def text(self, text, output_file): - request_json = { - "text": text, - "spk_id": self.spk_id, - "speed": self.speed, - "volume": self.volume, - "sample_rate": self.sample_rate, - "save_path": self.save_path - } - - try: - async with aiohttp.ClientSession() as session: - async with session.post(self.url, json=request_json) as resp: - if resp.status == 200: - resp_json = await resp.json() - if resp_json.get("success"): - data = resp_json["result"] - audio_bytes = base64.b64decode(data["audio"]) - if output_file: - with open(output_file, "wb") as file_to_save: - file_to_save.write(audio_bytes) - else: - return audio_bytes - else: - raise Exception( - f"Error: {resp_json.get('message', 'Unknown error')} while processing text: {text}") - else: - raise Exception( - f"HTTP Error: {resp.status} - {await resp.text()} while processing text: {text}") - except Exception as e: - raise Exception(f"Error during TTS HTTP request: {e} while processing text: {text}") - async def text_streaming(self, text, output_file): try: # 使用 websockets 异步连接到 WebSocket 服务器 @@ -151,6 +134,12 @@ class TTSProvider(TTSProviderBase): # 接收结束响应避免服务抛出异常 await ws.recv() + # 根据配置决定是否保存文件 + if not self.delete_audio_file and self.save_path: + with open(self.save_path, "wb") as f: + f.write(wav_data) + logger.bind(tag=TAG).info(f"音频文件已保存到: {self.save_path}") + # 返回或保存音频数据 if output_file: with open(output_file, "wb") as file_to_save: @@ -159,4 +148,4 @@ class TTSProvider(TTSProviderBase): return wav_data except Exception as e: - raise Exception(f"Error during TTS WebSocket request: {e} while processing text: {text}") + raise Exception(f"Error during TTS WebSocket request: {e} while processing text: {text}") \ No newline at end of file diff --git a/main/xiaozhi-server/core/providers/vad/silero.py b/main/xiaozhi-server/core/providers/vad/silero.py index 95ab8ff3..b516d8fb 100644 --- a/main/xiaozhi-server/core/providers/vad/silero.py +++ b/main/xiaozhi-server/core/providers/vad/silero.py @@ -33,8 +33,8 @@ class VADProvider(VADProviderBase): int(min_silence_duration_ms) if min_silence_duration_ms else 1000 ) - # 至少要多少帧才算有语音,增加灵敏度 - self.frame_window_threshold = 1 + # 至少要多少帧才算有语音 + self.frame_window_threshold = 3 def is_vad(self, conn, opus_packet): try: diff --git a/main/xiaozhi-server/core/utils/opus_encoder_utils.py b/main/xiaozhi-server/core/utils/opus_encoder_utils.py index 8ca406d3..ae7066ce 100644 --- a/main/xiaozhi-server/core/utils/opus_encoder_utils.py +++ b/main/xiaozhi-server/core/utils/opus_encoder_utils.py @@ -6,10 +6,9 @@ Opus编码工具类 import logging import traceback import numpy as np -from typing import Optional, Callable, Any from opuslib_next import Encoder from opuslib_next import constants - +from typing import Optional, Callable, Any class OpusEncoderUtils: """PCM到Opus的编码器""" @@ -130,4 +129,4 @@ class OpusEncoderUtils: def close(self): """关闭编码器并释放资源""" # opuslib没有明确的关闭方法,Python的垃圾回收会处理 - pass + pass \ No newline at end of file diff --git a/main/xiaozhi-server/core/utils/p3.py b/main/xiaozhi-server/core/utils/p3.py index 415e5366..c75b968e 100644 --- a/main/xiaozhi-server/core/utils/p3.py +++ b/main/xiaozhi-server/core/utils/p3.py @@ -1,12 +1,15 @@ -import io import struct -from typing import Callable, Any +def decode_opus_from_file(input_file): + """ + 从p3文件中解码 Opus 数据,并返回一个 Opus 数据包的列表以及总时长。 + """ + opus_datas = [] + total_frames = 0 + sample_rate = 16000 # 文件采样率 + frame_duration_ms = 60 # 帧时长 + frame_size = int(sample_rate * frame_duration_ms / 1000) -def decode_opus_from_file_stream(input_file, callback: Callable[[Any], Any]): - """ - 从p3文件中解码 Opus 数据,由 callback 处理 Opus 数据包。 - """ with open(input_file, 'rb') as f: while True: # 读取头部(4字节):[1字节类型,1字节保留,2字节长度] @@ -22,13 +25,23 @@ def decode_opus_from_file_stream(input_file, callback: Callable[[Any], Any]): if len(opus_data) != data_len: raise ValueError(f"Data length({len(opus_data)}) mismatch({data_len}) in the file.") - callback(opus_data) + opus_datas.append(opus_data) + total_frames += 1 + # 计算总时长 + total_duration = (total_frames * frame_duration_ms) / 1000.0 + return opus_datas, total_duration -def decode_opus_from_bytes_stream(input_bytes, callback: Callable[[Any], Any]): +def decode_opus_from_bytes(input_bytes): """ - 从p3二进制数据中解码 Opus 数据,由 callback 处理 Opus 数据包。 + 从p3二进制数据中解码 Opus 数据,并返回一个 Opus 数据包的列表以及总时长。 """ + import io + opus_datas = [] + total_frames = 0 + sample_rate = 16000 # 文件采样率 + frame_duration_ms = 60 # 帧时长 + frame_size = int(sample_rate * frame_duration_ms / 1000) f = io.BytesIO(input_bytes) while True: @@ -39,4 +52,8 @@ def decode_opus_from_bytes_stream(input_bytes, callback: Callable[[Any], Any]): opus_data = f.read(data_len) if len(opus_data) != data_len: raise ValueError(f"Data length({len(opus_data)}) mismatch({data_len}) in the bytes.") - callback(opus_data) + opus_datas.append(opus_data) + total_frames += 1 + + total_duration = (total_frames * frame_duration_ms) / 1000.0 + return opus_datas, total_duration \ No newline at end of file diff --git a/main/xiaozhi-server/core/utils/util.py b/main/xiaozhi-server/core/utils/util.py index 1eec7327..27471e58 100644 --- a/main/xiaozhi-server/core/utils/util.py +++ b/main/xiaozhi-server/core/utils/util.py @@ -1,16 +1,17 @@ -import json -import socket -import subprocess import re import os -from io import BytesIO -from typing import Callable, Any -from core.utils import p3 -import numpy as np -import requests -import opuslib_next -from pydub import AudioSegment +import json import copy +import wave +import socket +import requests +import subprocess +import numpy as np +import opuslib_next +from io import BytesIO +from core.utils import p3 +from pydub import AudioSegment +from typing import Callable, Any TAG = __name__ emoji_map = { @@ -228,6 +229,56 @@ def audio_to_data_stream(audio_file_path, is_opus=True, callback: Callable[[Any] raw_data = audio.raw_data pcm_to_data_stream(raw_data, is_opus, callback) +def audio_to_data(audio_file_path: str, is_opus: bool = True) -> list[bytes]: + """ + 将音频文件转换为Opus/PCM编码的帧列表 + Args: + audio_file_path: 音频文件路径 + is_opus: 是否进行Opus编码 + """ + # 获取文件后缀名 + 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"] + ) + + # 转换为单声道/16kHz采样率/16位小端编码(确保与编码器匹配) + audio = audio.set_channels(1).set_frame_rate(16000).set_sample_width(2) + + # 获取原始PCM数据(16位小端) + raw_data = audio.raw_data + + # 初始化Opus编码器 + encoder = opuslib_next.Encoder(16000, 1, opuslib_next.APPLICATION_AUDIO) + + # 编码参数 + frame_duration = 60 # 60ms per frame + frame_size = int(16000 * frame_duration / 1000) # 960 samples/frame + + datas = [] + # 按帧处理所有音频数据(包括最后一帧可能补零) + for i in range(0, len(raw_data), frame_size * 2): # 16bit=2bytes/sample + # 获取当前帧的二进制数据 + chunk = raw_data[i : i + frame_size * 2] + + # 如果最后一帧不足,补零 + 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 def audio_bytes_to_data_stream(audio_bytes, file_type, is_opus, callback: Callable[[Any], Any]) -> None: """ @@ -273,6 +324,31 @@ def pcm_to_data_stream(raw_data, is_opus=True, callback: Callable[[Any], Any] = frame_data = chunk if isinstance(chunk, bytes) else bytes(chunk) callback(frame_data) +def opus_datas_to_wav_bytes(opus_datas, sample_rate=16000, channels=1): + """ + 将opus帧列表解码为wav字节流 + """ + decoder = opuslib_next.Decoder(sample_rate, channels) + pcm_datas = [] + + 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) + + 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() def check_vad_update(before_config, new_config): if ( diff --git a/main/xiaozhi-server/core/utils/wakeup_word.py b/main/xiaozhi-server/core/utils/wakeup_word.py new file mode 100644 index 00000000..d2f4fb32 --- /dev/null +++ b/main/xiaozhi-server/core/utils/wakeup_word.py @@ -0,0 +1,140 @@ +import os +import re +import yaml +import time +import hashlib +import portalocker +from typing import Dict + + +class FileLock: + def __init__(self, file, timeout=5): + self.file = file + self.timeout = timeout + self.start_time = None + + def __enter__(self): + self.start_time = time.time() + while True: + try: + portalocker.lock(self.file, portalocker.LOCK_EX | portalocker.LOCK_NB) + return self.file + except portalocker.LockException: + if time.time() - self.start_time > self.timeout: + raise TimeoutError("获取文件锁超时") + time.sleep(0.1) + + def __exit__(self, exc_type, exc_val, exc_tb): + portalocker.unlock(self.file) + + +class WakeupWordsConfig: + def __init__(self): + self.config_file = "data/.wakeup_words.yaml" + self.assets_dir = "config/assets/wakeup_words" + self._ensure_directories() + self._config_cache = None + self._last_load_time = 0 + self._cache_ttl = 1 # 缓存有效期(秒) + self._lock_timeout = 5 # 文件锁超时时间(秒) + + def _ensure_directories(self): + """确保必要的目录存在""" + os.makedirs(os.path.dirname(self.config_file), exist_ok=True) + os.makedirs(self.assets_dir, exist_ok=True) + + def _load_config(self) -> Dict: + """加载配置文件,使用缓存机制""" + current_time = time.time() + + # 如果缓存有效,直接返回缓存 + if ( + self._config_cache is not None + and current_time - self._last_load_time < self._cache_ttl + ): + return self._config_cache + + try: + with open(self.config_file, "a+") as f: + with FileLock(f, timeout=self._lock_timeout): + f.seek(0) + content = f.read() + config = yaml.safe_load(content) if content else {} + self._config_cache = config + self._last_load_time = current_time + return config + except (TimeoutError, IOError) as e: + print(f"加载配置文件失败: {e}") + return {} + except Exception as e: + print(f"加载配置文件时发生未知错误: {e}") + return {} + + def _save_config(self, config: Dict): + """保存配置到文件,使用文件锁保护""" + try: + with open(self.config_file, "w") as f: + with FileLock(f, timeout=self._lock_timeout): + yaml.dump(config, f, allow_unicode=True) + self._config_cache = config + self._last_load_time = time.time() + except (TimeoutError, IOError) as e: + print(f"保存配置文件失败: {e}") + raise + except Exception as e: + print(f"保存配置文件时发生未知错误: {e}") + raise + + def get_wakeup_response(self, voice: str) -> Dict: + voice = hashlib.md5(voice.encode()).hexdigest() + """获取唤醒词回复配置""" + config = self._load_config() + + if not config or voice not in config: + return None + + # 检查文件大小 + file_path = config[voice]["file_path"] + if not os.path.exists(file_path) or os.stat(file_path).st_size < (15 * 1024): + return None + + return config[voice] + + def update_wakeup_response(self, voice: str, file_path: str, text: str): + """更新唤醒词回复配置""" + try: + # 过滤表情符号 + filtered_text = re.sub(r'[\U0001F600-\U0001F64F\U0001F900-\U0001F9FF]', '', text) + + config = self._load_config() + voice_hash = hashlib.md5(voice.encode()).hexdigest() + config[voice_hash] = { + "voice": voice, + "file_path": file_path, + "time": time.time(), + "text": filtered_text, + } + self._save_config(config) + except Exception as e: + print(f"更新唤醒词回复配置失败: {e}") + raise + + def generate_file_path(self, voice: str) -> str: + """生成音频文件路径,使用voice的哈希值作为文件名""" + try: + # 生成voice的哈希值 + voice_hash = hashlib.md5(voice.encode()).hexdigest() + file_path = os.path.join(self.assets_dir, f"{voice_hash}.wav") + + # 如果文件已存在,先删除 + if os.path.exists(file_path): + try: + os.remove(file_path) + except Exception as e: + print(f"删除已存在的音频文件失败: {e}") + raise + + return file_path + except Exception as e: + print(f"生成音频文件路径失败: {e}") + raise \ No newline at end of file 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 54641106..1d60aefd 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 @@ -125,7 +125,8 @@ def fetch_news_from_api(conn, source="thepaper"): ]["get_news_from_newsnow"].get("url"): api_url = conn.config["plugins"]["get_news_from_newsnow"]["url"] + source - response = requests.get(api_url, timeout=10) + headers = {"User-Agent": "Mozilla/5.0"} + response = requests.get(api_url, headers=headers, timeout=10) response.raise_for_status() data = response.json() @@ -144,7 +145,8 @@ def fetch_news_from_api(conn, source="thepaper"): def fetch_news_detail(url): """获取新闻详情页内容并使用MarkItDown清理HTML""" try: - response = requests.get(url, timeout=10) + headers = {"User-Agent": "Mozilla/5.0"} + response = requests.get(url, headers=headers, timeout=10) response.raise_for_status() # 使用MarkItDown清理HTML内容 diff --git a/main/xiaozhi-server/plugins_func/functions/hass_get_state.py b/main/xiaozhi-server/plugins_func/functions/hass_get_state.py index 94478e35..4efb12b8 100644 --- a/main/xiaozhi-server/plugins_func/functions/hass_get_state.py +++ b/main/xiaozhi-server/plugins_func/functions/hass_get_state.py @@ -29,12 +29,7 @@ hass_get_state_function_desc = { @register_function("hass_get_state", hass_get_state_function_desc, ToolType.SYSTEM_CTL) def hass_get_state(conn, entity_id=""): try: - - future = asyncio.run_coroutine_threadsafe( - handle_hass_get_state(conn, entity_id), conn.loop - ) - # 添加10秒超时 - ha_response = future.result(timeout=10) + ha_response = handle_hass_get_state(conn, entity_id) return ActionResponse(Action.REQLLM, ha_response, None) except asyncio.TimeoutError: logger.bind(tag=TAG).error("获取Home Assistant状态超时") @@ -45,13 +40,13 @@ def hass_get_state(conn, entity_id=""): return ActionResponse(Action.ERROR, error_msg, None) -async def handle_hass_get_state(conn, entity_id): +def handle_hass_get_state(conn, entity_id): ha_config = initialize_hass_handler(conn) api_key = ha_config.get("api_key") base_url = ha_config.get("base_url") url = f"{base_url}/api/states/{entity_id}" headers = {"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"} - response = requests.get(url, headers=headers) + response = requests.get(url, headers=headers, timeout=5) if response.status_code == 200: responsetext = "设备状态:" + response.json()["state"] + " " logger.bind(tag=TAG).info(f"api返回内容: {response.json()}") diff --git a/main/xiaozhi-server/plugins_func/functions/hass_set_state.py b/main/xiaozhi-server/plugins_func/functions/hass_set_state.py index c1fcdaa6..3addc823 100644 --- a/main/xiaozhi-server/plugins_func/functions/hass_set_state.py +++ b/main/xiaozhi-server/plugins_func/functions/hass_set_state.py @@ -54,11 +54,7 @@ def hass_set_state(conn, entity_id="", state=None): if state is None: state = {} try: - future = asyncio.run_coroutine_threadsafe( - handle_hass_set_state(conn, entity_id, state), conn.loop - ) - # 添加10秒超时 - ha_response = future.result(timeout=10) + ha_response = handle_hass_set_state(conn, entity_id, state) return ActionResponse(Action.REQLLM, ha_response, None) except asyncio.TimeoutError: logger.bind(tag=TAG).error("设置Home Assistant状态超时") @@ -69,7 +65,7 @@ def hass_set_state(conn, entity_id="", state=None): return ActionResponse(Action.ERROR, error_msg, None) -async def handle_hass_set_state(conn, entity_id, state): +def handle_hass_set_state(conn, entity_id, state): ha_config = initialize_hass_handler(conn) api_key = ha_config.get("api_key") base_url = ha_config.get("base_url") @@ -169,7 +165,7 @@ async def handle_hass_set_state(conn, entity_id, state): data = {"entity_id": entity_id, arg: value} url = f"{base_url}/api/services/{domain}/{action}" headers = {"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"} - response = requests.post(url, headers=headers, json=data) + response = requests.post(url, headers=headers, json=data, timeout=5) # 设置5秒超时 logger.bind(tag=TAG).info( f"设置状态:{description},url:{url},return_code:{response.status_code}" ) diff --git a/main/xiaozhi-server/plugins_func/functions/play_music.py b/main/xiaozhi-server/plugins_func/functions/play_music.py index dd967d0a..2cbc4018 100644 --- a/main/xiaozhi-server/plugins_func/functions/play_music.py +++ b/main/xiaozhi-server/plugins_func/functions/play_music.py @@ -212,6 +212,7 @@ async def play_local_music(conn, specific_file=None): conn.logger.bind(tag=TAG).error(f"选定的音乐文件不存在: {music_path}") return text = _get_random_play_prompt(selected_music) + await send_stt_message(conn, text) conn.dialogue.put(Message(role="assistant", content=text)) if conn.intent_type == "intent_llm": diff --git a/main/xiaozhi-server/test/js/StreamingContext.js b/main/xiaozhi-server/test/js/StreamingContext.js new file mode 100644 index 00000000..576d9761 --- /dev/null +++ b/main/xiaozhi-server/test/js/StreamingContext.js @@ -0,0 +1,149 @@ +import BlockingQueue from './utils/BlockingQueue.js'; +import { log } from './utils/logger.js'; + +// 音频流播放上下文类 +export class StreamingContext { + constructor(opusDecoder, audioContext, sampleRate, channels, minAudioDuration) { + this.opusDecoder = opusDecoder; + this.audioContext = audioContext; + + // 音频参数 + this.sampleRate = sampleRate; + this.channels = channels; + this.minAudioDuration = minAudioDuration; + + // 初始化队列和状态 + this.queue = []; // 已解码的PCM队列。正在播放 + this.activeQueue = new BlockingQueue(); // 已解码的PCM队列。准备播放 + this.pendingAudioBufferQueue = []; // 待处理的缓存队列 + this.audioBufferQueue = new BlockingQueue(); // 缓存队列 + this.playing = false; // 是否正在播放 + this.endOfStream = false; // 是否收到结束信号 + this.source = null; // 当前音频源 + this.totalSamples = 0; // 累积的总样本数 + this.lastPlayTime = 0; // 上次播放的时间戳 + } + + // 缓存音频数组 + pushAudioBuffer(item) { + this.audioBufferQueue.enqueue(...item); + } + + // 获取需要处理缓存队列,单线程:在audioBufferQueue一直更新的状态下不会出现安全问题 + async getPendingAudioBufferQueue() { + // 原子交换 + 清空 + [this.pendingAudioBufferQueue, this.audioBufferQueue] = [await this.audioBufferQueue.dequeue(), new BlockingQueue()]; + } + + // 获取正在播放已解码的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); + } + + // 将Int16音频数据转换为Float32音频数据 + convertInt16ToFloat32(int16Data) { + const float32Data = new Float32Array(int16Data.length); + for (let i = 0; i < int16Data.length; i++) { + // 将[-32768,32767]范围转换为[-1,1] + float32Data[i] = int16Data[i] / (int16Data[i] < 0 ? 0x8000 : 0x7FFF); + } + return float32Data; + } + + // 将Opus数据解码为PCM + async decodeOpusFrames() { + if (!this.opusDecoder) { + log('Opus解码器未初始化,无法解码', 'error'); + return; + } else { + log('Opus解码器启动', 'info'); + } + + while (true) { + let decodedSamples = []; + for (const frame of this.pendingAudioBufferQueue) { + try { + // 使用Opus解码器解码 + const frameData = this.opusDecoder.decode(frame); + if (frameData && frameData.length > 0) { + // 转换为Float32 + const floatData = this.convertInt16ToFloat32(frameData); + // 使用循环替代展开运算符 + for (let i = 0; i < floatData.length; i++) { + decodedSamples.push(floatData[i]); + } + } + } catch (error) { + log("Opus解码失败: " + error.message, 'error'); + } + } + + if (decodedSamples.length > 0) { + // 使用循环替代展开运算符 + for (let i = 0; i < decodedSamples.length; i++) { + this.activeQueue.enqueue(decodedSamples[i]); + } + this.totalSamples += decodedSamples.length; + } else { + log('没有成功解码的样本', 'warning'); + } + await this.getPendingAudioBufferQueue(); + } + } + + // 开始播放音频 + async startPlaying() { + while (true) { + // 如果累积了至少0.3秒的音频,开始播放 + const minSamples = this.sampleRate * this.minAudioDuration * 3; + if (!this.playing && this.queue.length < minSamples) { + await this.getQueue(minSamples); + } + this.playing = true; + while (this.playing && this.queue.length) { + // 创建新的音频缓冲区 + const minPlaySamples = Math.min(this.queue.length, this.sampleRate); + const currentSamples = this.queue.splice(0, minPlaySamples); + + const audioBuffer = this.audioContext.createBuffer(this.channels, currentSamples.length, this.sampleRate); + audioBuffer.copyToChannel(new Float32Array(currentSamples), 0); + + // 创建音频源 + this.source = this.audioContext.createBufferSource(); + this.source.buffer = audioBuffer; + + // 创建增益节点用于平滑过渡 + const gainNode = this.audioContext.createGain(); + + // 应用淡入淡出效果避免爆音 + const fadeDuration = 0.02; // 20毫秒 + gainNode.gain.setValueAtTime(0, this.audioContext.currentTime); + gainNode.gain.linearRampToValueAtTime(1, this.audioContext.currentTime + fadeDuration); + + const duration = audioBuffer.duration; + if (duration > fadeDuration * 2) { + gainNode.gain.setValueAtTime(1, this.audioContext.currentTime + duration - fadeDuration); + gainNode.gain.linearRampToValueAtTime(0, this.audioContext.currentTime + duration); + } + + // 连接节点并开始播放 + this.source.connect(gainNode); + gainNode.connect(this.audioContext.destination); + + this.lastPlayTime = this.audioContext.currentTime; + log(`开始播放 ${currentSamples.length} 个样本,约 ${(currentSamples.length / this.sampleRate).toFixed(2)} 秒`, 'info'); + this.source.start(); + } + await this.getQueue(minSamples); + } + } +} + +// 创建streamingContext实例的工厂函数 +export function createStreamingContext(opusDecoder, audioContext, sampleRate, channels, minAudioDuration) { + return new StreamingContext(opusDecoder, audioContext, sampleRate, channels, minAudioDuration); +} \ 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 983d3b9a..547f2b31 100644 --- a/main/xiaozhi-server/test/test_page.html +++ b/main/xiaozhi-server/test/test_page.html @@ -181,6 +181,7 @@ import { checkOpusLoaded, initOpusEncoder } from './js/opus.js'; import { addMessage } from './js/document.js' import BlockingQueue from './js/utils/BlockingQueue.js' + import { createStreamingContext } from './js/StreamingContext.js' // 需要加载的脚本列表 - 移除Opus依赖 const scriptFiles = []; @@ -336,125 +337,7 @@ // 创建流式播放上下文 if (!streamingContext) { - streamingContext = { - queue: [], // 已解码的PCM队列。正在播放 - activeQueue: new BlockingQueue(), // 已解码的PCM队列。准备播放 - pendingAudioBufferQueue: [], // 待处理的缓存队列 - audioBufferQueue: new BlockingQueue(), // 缓存队列 - playing: false, // 是否正在播放 - endOfStream: false, // 是否收到结束信号 - source: null, // 当前音频源 - totalSamples: 0, // 累积的总样本数 - lastPlayTime: 0, // 上次播放的时间戳 - - - // 缓存音频数组 - pushAudioBuffer: function (item) { - this.audioBufferQueue.enqueue(...item) - }, - - // 获取需要处理缓存队列,单线程:在audioBufferQueue一直更新的状态下不会出现安全问题 - getPendingAudioBufferQueue: async function () { - // 原子交换 + 清空 - [this.pendingAudioBufferQueue, this.audioBufferQueue] = [await this.audioBufferQueue.dequeue(), new BlockingQueue()]; - - }, - // 获取正在播放已解码的PCM队列,单线程:在activeQueue一直更新的状态下不会出现安全问题 - getQueue: async function (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) - }, - // 将Opus数据解码为PCM - decodeOpusFrames: async function () { - if (!opusDecoder) { - log('Opus解码器未初始化,无法解码', 'error'); - return; - } else { - log('Opus解码器启动', 'info'); - } - - while (true) { - let decodedSamples = []; - for (const frame of this.pendingAudioBufferQueue) { - try { - // 使用Opus解码器解码 - const frameData = opusDecoder.decode(frame); - if (frameData && frameData.length > 0) { - // 转换为Float32 - const floatData = convertInt16ToFloat32(frameData); - // 使用循环替代展开运算符 - for (let i = 0; i < floatData.length; i++) { - decodedSamples.push(floatData[i]); - } - } - } catch (error) { - log("Opus解码失败: " + error.message, 'error'); - } - } - - if (decodedSamples.length > 0) { - // 使用循环替代展开运算符 - for (let i = 0; i < decodedSamples.length; i++) { - this.activeQueue.enqueue(decodedSamples[i]); - } - this.totalSamples += decodedSamples.length; - } else { - log('没有成功解码的样本', 'warning'); - } - await this.getPendingAudioBufferQueue(); - } - }, - - // 开始播放音频 - startPlaying: async function () { - while (true) { - // 如果累积了至少0.3秒的音频,开始播放 - const minSamples = SAMPLE_RATE * MIN_AUDIO_DURATION * 3; - if (!this.playing && this.queue.length < minSamples) { - await this.getQueue(minSamples) - } - this.playing = true; - while (this.playing && this.queue.length) { - // 创建新的音频缓冲区 - const minPlaySamples = Math.min(this.queue.length, SAMPLE_RATE); - const currentSamples = this.queue.splice(0, minPlaySamples); - - const audioBuffer = audioContext.createBuffer(CHANNELS, currentSamples.length, SAMPLE_RATE); - audioBuffer.copyToChannel(new Float32Array(currentSamples), 0); - - // 创建音频源 - this.source = audioContext.createBufferSource(); - this.source.buffer = audioBuffer; - - // 创建增益节点用于平滑过渡 - const gainNode = audioContext.createGain(); - - // 应用淡入淡出效果避免爆音 - const fadeDuration = 0.02; // 20毫秒 - gainNode.gain.setValueAtTime(0, audioContext.currentTime); - gainNode.gain.linearRampToValueAtTime(1, audioContext.currentTime + fadeDuration); - - const duration = audioBuffer.duration; - if (duration > fadeDuration * 2) { - gainNode.gain.setValueAtTime(1, audioContext.currentTime + duration - fadeDuration); - gainNode.gain.linearRampToValueAtTime(0, audioContext.currentTime + duration); - } - - // 连接节点并开始播放 - this.source.connect(gainNode); - gainNode.connect(audioContext.destination); - - this.lastPlayTime = audioContext.currentTime; - log(`开始播放 ${currentSamples.length} 个样本,约 ${(currentSamples.length / SAMPLE_RATE).toFixed(2)} 秒`, 'info'); - this.source.start(); - } - await this.getQueue(minSamples) - } - } - }; + streamingContext = createStreamingContext(opusDecoder, audioContext, SAMPLE_RATE, CHANNELS, MIN_AUDIO_DURATION); } streamingContext.decodeOpusFrames(); @@ -467,15 +350,7 @@ } } - // 将Int16音频数据转换为Float32音频数据 - function convertInt16ToFloat32(int16Data) { - const float32Data = new Float32Array(int16Data.length); - for (let i = 0; i < int16Data.length; i++) { - // 将[-32768,32767]范围转换为[-1,1] - float32Data[i] = int16Data[i] / (int16Data[i] < 0 ? 0x8000 : 0x7FFF); - } - return float32Data; - } + // 初始化Opus解码器 - 确保完全初始化完成后才返回 async function initOpusDecoder() {