From 610fa4d10185fd6d122df9e27888a1735506faf7 Mon Sep 17 00:00:00 2001
From: hrz <1710360675@qq.com>
Date: Tue, 3 Jun 2025 17:30:35 +0800
Subject: [PATCH] =?UTF-8?q?update:=E5=8C=BA=E5=88=86=E8=B1=86=E5=8C=85ASR?=
=?UTF-8?q?=E6=8C=89=E6=AC=A1=E6=94=B6=E8=B4=B9=E5=92=8C=E6=8C=89=E6=97=B6?=
=?UTF-8?q?=E6=94=B6=E8=B4=B9=E6=8E=A5=E5=8F=A3?=
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit
---
.../xiaozhi/common/constant/Constant.java | 2 +-
.../resources/db/changelog/202506031639.sql | 45 ++
main/manager-web/src/views/ModelConfig.vue | 4 +-
main/manager-web/src/views/roleConfig.vue | 2 +-
main/xiaozhi-server/config.yaml | 15 +
main/xiaozhi-server/config/logger.py | 2 +-
.../core/providers/asr/doubao.py | 679 ++++++------------
.../core/providers/asr/doubao_stream.py | 534 ++++++++++++++
.../xiaozhi-server/core/providers/tts/base.py | 1 +
9 files changed, 806 insertions(+), 478 deletions(-)
create mode 100644 main/manager-api/src/main/resources/db/changelog/202506031639.sql
create mode 100644 main/xiaozhi-server/core/providers/asr/doubao_stream.py
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 2e3174af..2590a566 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
@@ -227,7 +227,7 @@ public interface Constant {
/**
* 版本号
*/
- public static final String VERSION = "0.5.2";
+ public static final String VERSION = "0.5.4";
/**
* 无效固件URL
diff --git a/main/manager-api/src/main/resources/db/changelog/202506031639.sql b/main/manager-api/src/main/resources/db/changelog/202506031639.sql
new file mode 100644
index 00000000..95691b42
--- /dev/null
+++ b/main/manager-api/src/main/resources/db/changelog/202506031639.sql
@@ -0,0 +1,45 @@
+-- VLLM模型供应器
+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"}]', 3, 1, NOW(), 1, NOW());
+
+
+-- VLLM模型配置
+delete from `ai_model_config` where id = 'ASR_DoubaoStreamASR';
+INSERT INTO `ai_model_config` VALUES ('ASR_DoubaoStreamASR', 'ASR', 'DoubaoStreamASR', '豆包语音识别(流式)', 0, 1, '{\"type\": \"doubao_stream\", \"appid\": \"\", \"access_token\": \"\", \"cluster\": \"volcengine_input_common\", \"output_dir\": \"tmp/\"}', NULL, NULL, 3, NULL, NULL, NULL, NULL);
+
+
+-- 更新豆包ASR配置说明
+UPDATE `ai_model_config` SET
+`doc_link` = 'https://console.volcengine.com/speech/app',
+`remark` = '豆包ASR配置说明:
+1. 豆包ASR和豆包(流式)ASR的区别是:豆包ASR是按次收费,豆包(流式)ASR是按时收费
+2. 一般来说按次收费的更便宜,但是豆包(流式)ASR使用了大模型技术,效果更好
+3. 需要在火山引擎控制台创建应用并获取appid和access_token
+4. 支持中文语音识别
+5. 需要网络连接
+6. 输出文件保存在tmp/目录
+申请步骤:
+1. 访问 https://console.volcengine.com/speech/app
+2. 创建新应用
+3. 获取appid和access_token
+4. 填入配置文件中
+如需设置热词,请参考:https://www.volcengine.com/docs/6561/155738
+' WHERE `id` = 'ASR_DoubaoASR';
+
+UPDATE `ai_model_config` SET
+`doc_link` = 'https://console.volcengine.com/speech/app',
+`remark` = '豆包ASR配置说明:
+1. 豆包ASR和豆包(流式)ASR的区别是:豆包ASR是按次收费,豆包(流式)ASR是按时收费
+2. 一般来说按次收费的更便宜,但是豆包(流式)ASR使用了大模型技术,效果更好
+3. 需要在火山引擎控制台创建应用并获取appid和access_token
+4. 支持中文语音识别
+5. 需要网络连接
+6. 输出文件保存在tmp/目录
+申请步骤:
+1. 访问 https://console.volcengine.com/speech/app
+2. 创建新应用
+3. 获取appid和access_token
+4. 填入配置文件中
+如需设置热词,请参考:https://www.volcengine.com/docs/6561/155738
+' WHERE `id` = 'ASR_DoubaoStreamASR';
diff --git a/main/manager-web/src/views/ModelConfig.vue b/main/manager-web/src/views/ModelConfig.vue
index 6e5a76f8..5a88c9aa 100644
--- a/main/manager-web/src/views/ModelConfig.vue
+++ b/main/manager-web/src/views/ModelConfig.vue
@@ -31,7 +31,7 @@
-
+
@@ -176,7 +176,7 @@ export default {
vad: '语言活动检测模型(VAD)',
asr: '语音识别模型(ASR)',
llm: '大语言模型(LLM)',
- vllm: '视觉大语言模型(VLLM)',
+ vllm: '视觉大模型(VLLM)',
intent: '意图识别模型(Intent)',
tts: '语音合成模型(TTS)',
memory: '记忆模型(Memory)'
diff --git a/main/manager-web/src/views/roleConfig.vue b/main/manager-web/src/views/roleConfig.vue
index 50cf278a..06a5c09d 100644
--- a/main/manager-web/src/views/roleConfig.vue
+++ b/main/manager-web/src/views/roleConfig.vue
@@ -177,7 +177,7 @@ export default {
{ label: '语音活动检测(VAD)', key: 'vadModelId', type: 'VAD' },
{ label: '语音识别(ASR)', key: 'asrModelId', type: 'ASR' },
{ label: '大语言模型(LLM)', key: 'llmModelId', type: 'LLM' },
- { label: '视觉大语言模型(VLLM)', key: 'vllmModelId', type: 'VLLM' },
+ { label: '视觉大模型(VLLM)', key: 'vllmModelId', type: 'VLLM' },
{ label: '意图识别(Intent)', key: 'intentModelId', type: 'Intent' },
{ label: '记忆(Memory)', key: 'memModelId', type: 'Memory' },
{ label: '语音合成(TTS)', key: 'ttsModelId', type: 'TTS' },
diff --git a/main/xiaozhi-server/config.yaml b/main/xiaozhi-server/config.yaml
index 37a7309d..179d7204 100644
--- a/main/xiaozhi-server/config.yaml
+++ b/main/xiaozhi-server/config.yaml
@@ -264,6 +264,8 @@ ASR:
DoubaoASR:
# 可以在这里申请相关Key等信息
# https://console.volcengine.com/speech/app
+ # DoubaoASR和DoubaoStreamASR的区别是:DoubaoASR是按次收费,DoubaoStreamASR是按时收费
+ # 一般来说按次收费的更便宜,但是DoubaoStreamASR使用了大模型技术,效果更好
type: doubao
appid: 你的火山引擎语音合成服务appid
access_token: 你的火山引擎语音合成服务access_token
@@ -272,6 +274,19 @@ ASR:
boosting_table_name: (选填)你的热词文件名称
correct_table_name: (选填)你的替换词文件名称
output_dir: tmp/
+ DoubaoStreamASR:
+ # 可以在这里申请相关Key等信息
+ # https://console.volcengine.com/speech/app
+ # DoubaoASR和DoubaoStreamASR的区别是:DoubaoASR是按次收费,DoubaoStreamASR是按时收费
+ # 一般来说按次收费的更便宜,但是DoubaoStreamASR使用了大模型技术,效果更好
+ type: doubao_stream
+ appid: 你的火山引擎语音合成服务appid
+ access_token: 你的火山引擎语音合成服务access_token
+ cluster: volcengine_input_common
+ # 热词、替换词使用流程:https://www.volcengine.com/docs/6561/155738
+ boosting_table_name: (选填)你的热词文件名称
+ correct_table_name: (选填)你的替换词文件名称
+ output_dir: tmp/
TencentASR:
# token申请地址:https://console.cloud.tencent.com/cam/capi
# 免费领取资源:https://console.cloud.tencent.com/asr/resourcebundle
diff --git a/main/xiaozhi-server/config/logger.py b/main/xiaozhi-server/config/logger.py
index 5f107f5d..643952dd 100644
--- a/main/xiaozhi-server/config/logger.py
+++ b/main/xiaozhi-server/config/logger.py
@@ -4,7 +4,7 @@ from loguru import logger
from config.config_loader import load_config
from config.settings import check_config_file
-SERVER_VERSION = "0.5.2"
+SERVER_VERSION = "0.5.4"
_logger_initialized = False
diff --git a/main/xiaozhi-server/core/providers/asr/doubao.py b/main/xiaozhi-server/core/providers/asr/doubao.py
index da4098ad..8a647508 100644
--- a/main/xiaozhi-server/core/providers/asr/doubao.py
+++ b/main/xiaozhi-server/core/providers/asr/doubao.py
@@ -1,534 +1,267 @@
+import time
+import os
+import uuid
import json
import gzip
-import uuid
-import asyncio
import websockets
-import opuslib_next
-from core.providers.asr.base import ASRProviderBase
from config.logger import setup_logging
+from typing import Optional, Tuple, List
+from core.providers.asr.base import ASRProviderBase
from core.providers.asr.dto.dto import InterfaceType
-import threading
+
TAG = __name__
logger = setup_logging()
CLIENT_FULL_REQUEST = 0b0001
CLIENT_AUDIO_ONLY_REQUEST = 0b0010
+
+NO_SEQUENCE = 0b0000
+NEG_SEQUENCE = 0b0010
+
SERVER_FULL_RESPONSE = 0b1001
SERVER_ACK = 0b1011
SERVER_ERROR_RESPONSE = 0b1111
-NO_SEQUENCE = 0b0000
-NEG_SEQUENCE = 0b0010
-JSON_SERIALIZATION = 0b0001
-GZIP_COMPRESSION = 0b0001
-PROTOCOL_VERSION = 0b0001
+
+NO_SERIALIZATION = 0b0000
+JSON = 0b0001
+THRIFT = 0b0011
+CUSTOM_TYPE = 0b1111
+NO_COMPRESSION = 0b0000
+GZIP = 0b0001
+CUSTOM_COMPRESSION = 0b1111
+
+
+def parse_response(res):
+ """
+ protocol_version(4 bits), header_size(4 bits),
+ message_type(4 bits), message_type_specific_flags(4 bits)
+ serialization_method(4 bits) message_compression(4 bits)
+ reserved (8bits) 保留字段
+ header_extensions 扩展头(大小等于 8 * 4 * (header_size - 1) )
+ payload 类似与http 请求体
+ """
+ protocol_version = res[0] >> 4
+ header_size = res[0] & 0x0F
+ message_type = res[1] >> 4
+ message_type_specific_flags = res[1] & 0x0F
+ serialization_method = res[2] >> 4
+ message_compression = res[2] & 0x0F
+ reserved = res[3]
+ header_extensions = res[4 : header_size * 4]
+ payload = res[header_size * 4 :]
+ result = {}
+ payload_msg = None
+ payload_size = 0
+ if message_type == SERVER_FULL_RESPONSE:
+ payload_size = int.from_bytes(payload[:4], "big", signed=True)
+ payload_msg = payload[4:]
+ elif message_type == SERVER_ACK:
+ seq = int.from_bytes(payload[:4], "big", signed=True)
+ result["seq"] = seq
+ if len(payload) >= 8:
+ payload_size = int.from_bytes(payload[4:8], "big", signed=False)
+ payload_msg = payload[8:]
+ elif message_type == SERVER_ERROR_RESPONSE:
+ code = int.from_bytes(payload[:4], "big", signed=False)
+ result["code"] = code
+ payload_size = int.from_bytes(payload[4:8], "big", signed=False)
+ payload_msg = payload[8:]
+ if payload_msg is None:
+ return result
+ if message_compression == GZIP:
+ payload_msg = gzip.decompress(payload_msg)
+ if serialization_method == JSON:
+ payload_msg = json.loads(str(payload_msg, "utf-8"))
+ elif serialization_method != NO_SERIALIZATION:
+ payload_msg = str(payload_msg, "utf-8")
+ result["payload_msg"] = payload_msg
+ result["payload_size"] = payload_size
+ return result
class ASRProvider(ASRProviderBase):
- def __init__(self, config, delete_audio_file):
+ def __init__(self, config: dict, delete_audio_file: bool):
super().__init__()
- self.interface_type = InterfaceType.STREAM
- self.config = config
- self.text = ""
- self.max_retries = 3
- self.retry_delay = 2 # 重试延迟秒数
- self.recv_lock = asyncio.Lock() # 添加接收锁
- self.reconnect_lock = asyncio.Lock() # 添加重连锁
- self.last_reconnect_time = 0 # 上次重连时间
- self.reconnect_cooldown = 1 # 增加重连冷却时间到10秒
- self.reconnect_count = 0 # 当前重连次数
- self.max_reconnect_count = 3 # 减少最大重连次数到3次
- self.asr_thread = None # ASR监听线程
- self.thread_lock = threading.Lock() # 线程管理锁
- self.is_reconnecting = False # 添加重连状态标志
-
- # 添加会话管理相关属性
- self._session_lock = asyncio.Lock() # 会话操作的并发锁
- self._current_session_id = None # 当前会话ID
- self._session_started = False # 会话是否已开始
- self._session_finished = False # 会话是否已结束
- self._session_close_event = asyncio.Event() # 添加会话关闭事件
-
- self.appid = str(config.get("appid"))
+ self.interface_type = InterfaceType.NON_STREAM
+ self.appid = config.get("appid")
self.cluster = config.get("cluster")
self.access_token = config.get("access_token")
self.boosting_table_name = config.get("boosting_table_name", "")
self.correct_table_name = config.get("correct_table_name", "")
- self.output_dir = config.get("output_dir", "temp/")
+ self.output_dir = config.get("output_dir")
self.delete_audio_file = delete_audio_file
- self.ws_url = "wss://openspeech.bytedance.com/api/v2/asr"
- self.uid = config.get("uid", "streaming_asr_service")
- self.workflow = config.get(
- "workflow", "audio_in,resample,partition,vad,fe,decode,itn,nlu_punctuate"
- )
- self.result_type = config.get("result_type", "single")
- self.format = config.get("format", "raw")
- self.codec = config.get("codec", "pcm")
- self.rate = config.get("sample_rate", 16000)
- self.language = config.get("language", "zh-CN")
- self.bits = config.get("bits", 16)
- self.channel = config.get("channel", 1)
- self.auth_method = config.get("auth_method", "token")
- self.secret = config.get("secret", "access_secret")
- self.decoder = opuslib_next.Decoder(16000, 1)
- self.asr_ws = None
- self.forward_task = None
- self.conn = None
+ self.host = "openspeech.bytedance.com"
+ self.ws_url = f"wss://{self.host}/api/v2/asr"
+ self.success_code = 1000
+ self.seg_duration = 15000
- ###################################################################################
- # 豆包流式ASR重写父类的方法--开始
- ###################################################################################
- async def open_audio_channels(self, conn):
- await super().open_audio_channels(conn)
+ # 确保输出目录存在
+ os.makedirs(self.output_dir, exist_ok=True)
- async with self._session_lock:
- # 如果正在重连,等待重连完成
- if self.is_reconnecting:
- logger.bind(tag=TAG).info("等待当前重连完成...")
- await self._session_close_event.wait()
- self._session_close_event.clear()
+ @staticmethod
+ def _generate_header(
+ message_type=CLIENT_FULL_REQUEST, message_type_specific_flags=NO_SEQUENCE
+ ) -> bytearray:
+ """Generate protocol header."""
+ header = bytearray()
+ header_size = 1
+ header.append((0b0001 << 4) | header_size) # Protocol version
+ header.append((message_type << 4) | message_type_specific_flags)
+ header.append((0b0001 << 4) | 0b0001) # JSON serialization & GZIP compression
+ header.append(0x00) # reserved
+ return header
- # 如果已有会话未结束,先关闭它
- if self._session_started and not self._session_finished:
- logger.bind(tag=TAG).warning(
- f"发现未关闭的会话 {self._current_session_id},正在关闭..."
- )
- if self.asr_ws is not None:
- try:
- await self.asr_ws.close()
- except Exception as e:
- logger.bind(tag=TAG).warning(f"关闭旧连接时发生错误: {e}")
- finally:
- self.asr_ws = None
- self._session_finished = True
- self._session_close_event.set()
-
- # 重置会话状态
- self._current_session_id = str(uuid.uuid4())
- self._session_started = True
- self._session_finished = False
- self.is_reconnecting = True
-
- try:
- retry_count = 0
- while retry_count < self.max_retries:
- try:
- headers = (
- self.token_auth() if self.auth_method == "token" else None
- )
- self.asr_ws = await websockets.connect(
- self.ws_url,
- additional_headers=headers,
- max_size=1000000000,
- ping_interval=None,
- ping_timeout=None,
- close_timeout=10,
- )
-
- # 发送初始化请求
- request_params = self.construct_request(
- self._current_session_id
- )
- try:
- payload_bytes = str.encode(json.dumps(request_params))
- payload_bytes = gzip.compress(payload_bytes)
- full_client_request = self.generate_header()
- full_client_request.extend(
- (len(payload_bytes)).to_bytes(4, "big")
- )
- full_client_request.extend(payload_bytes)
- await self.asr_ws.send(full_client_request)
- except Exception as e:
- logger.bind(tag=TAG).error(f"发送初始化请求失败: {e}")
- raise e
-
- # 等待初始化响应
- try:
- init_res = await self.asr_ws.recv()
- self.parse_response(init_res)
- except Exception as e:
- logger.bind(tag=TAG).error(f"ASR服务初始化失败: {e}")
- raise e
-
- # 启动接收ASR结果的异步任务
- with self.thread_lock:
- if (
- self.asr_thread is None
- or not self.asr_thread.is_alive()
- ):
- logger.bind(tag=TAG).info("创建新的ASR监听线程...")
- self.asr_thread = threading.Thread(
- target=self._start_monitor_asr_response_thread,
- daemon=True,
- )
- self.asr_thread.start()
- # 等待一小段时间确保线程启动
- await asyncio.sleep(0.1)
- if not self.asr_thread.is_alive():
- logger.bind(tag=TAG).error("ASR监听线程启动失败")
- raise Exception("ASR监听线程启动失败")
- logger.bind(tag=TAG).info("ASR监听线程已启动")
- return
-
- except websockets.exceptions.WebSocketException as e:
- retry_count += 1
- if retry_count < self.max_retries:
- logger.bind(tag=TAG).warning(
- f"WebSocket连接失败,正在进行第{retry_count}次重试: {e}"
- )
- await asyncio.sleep(self.retry_delay)
- else:
- logger.bind(tag=TAG).warning(
- f"WebSocket连接失败,已达到最大重试次数: {e}"
- )
- raise
- except Exception as e:
- logger.bind(tag=TAG).error(f"WebSocket连接发生未知错误: {e}")
- raise
- finally:
- self.is_reconnecting = False
- self._session_close_event.set()
-
- async def receive_audio(self, audio, _):
- if not isinstance(audio, bytes):
- return
-
- try:
- # 解码opus得到PCM数据
- pcm_frame = self.decoder.decode(audio, 960)
- payload = gzip.compress(pcm_frame)
- audio_request = bytearray(self.generate_audio_default_header())
- audio_request.extend(len(payload).to_bytes(4, "big"))
- audio_request.extend(payload)
- if self.asr_ws:
- await self.asr_ws.send(audio_request)
- except Exception as e:
- logger.bind(tag=TAG).debug(f"发送音频数据时发生错误: {e}")
-
- ###################################################################################
- # 豆包流式ASR重写父类的方法--结束
- ###################################################################################
-
- def construct_request(self, reqid):
- req = {
+ def _construct_request(self, reqid) -> dict:
+ """Construct the request payload."""
+ return {
"app": {
- "appid": self.appid,
+ "appid": f"{self.appid}",
"cluster": self.cluster,
"token": self.access_token,
},
- "user": {"uid": self.uid},
+ "user": {
+ "uid": str(uuid.uuid4()),
+ },
"request": {
"reqid": reqid,
- "workflow": self.workflow,
- "show_utterances": True,
- "result_type": self.result_type,
+ "show_utterances": False,
"sequence": 1,
"boosting_table_name": self.boosting_table_name,
"correct_table_name": self.correct_table_name,
},
"audio": {
- "format": self.format,
- "codec": self.codec,
- "rate": self.rate,
- "language": self.language,
- "bits": self.bits,
- "channel": self.channel,
+ "format": "raw",
+ "rate": 16000,
+ "language": "zh-CN",
+ "bits": 16,
+ "channel": 1,
+ "codec": "raw",
},
}
- return req
- def token_auth(self):
- return {"Authorization": f"Bearer; {self.access_token}"}
-
- def generate_header(
- self,
- version=PROTOCOL_VERSION,
- message_type=CLIENT_FULL_REQUEST,
- message_type_specific_flags=NO_SEQUENCE,
- serial_method=JSON_SERIALIZATION,
- compression_type=GZIP_COMPRESSION,
- reserved_data=0x00,
- extension_header: bytes = b"",
- ):
- """
- 生成协议头:
- - 第1字节:高4位:协议版本,低4位:头部大小(单位 4 字节)
- - 第2字节:高4位:消息类型,低4位:消息类型特定标志
- - 第3字节:高4位:序列化方式,低4位:压缩方式
- - 第4字节:保留字段
- - 后续:扩展头(如果有)
- """
- header = bytearray()
- 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_full_default_header(self):
- # full client request 默认头
- return self.generate_header(
- version=PROTOCOL_VERSION,
- message_type=CLIENT_FULL_REQUEST,
- message_type_specific_flags=NO_SEQUENCE,
- serial_method=JSON_SERIALIZATION,
- compression_type=GZIP_COMPRESSION,
- )
-
- def generate_audio_default_header(self):
- # 普通音频片段请求
- return self.generate_header(
- version=PROTOCOL_VERSION,
- message_type=CLIENT_AUDIO_ONLY_REQUEST,
- message_type_specific_flags=NO_SEQUENCE,
- serial_method=JSON_SERIALIZATION,
- compression_type=GZIP_COMPRESSION,
- )
-
- def generate_last_audio_default_header(self):
- # 最后一个音频片段标志
- return self.generate_header(
- version=PROTOCOL_VERSION,
- message_type=CLIENT_AUDIO_ONLY_REQUEST,
- message_type_specific_flags=NEG_SEQUENCE, # 用 NEG_SEQUENCE 表示结束
- serial_method=JSON_SERIALIZATION,
- compression_type=GZIP_COMPRESSION,
- )
-
- def _start_monitor_asr_response_thread(self):
- # 初始化链接
+ async def _send_request(
+ self, audio_data: List[bytes], segment_size: int
+ ) -> Optional[str]:
+ """Send request to Volcano ASR service."""
try:
- with self.thread_lock:
- if self.conn is None or self.conn.loop is None:
- logger.bind(tag=TAG).error(
- "无法启动ASR监听线程:conn或loop未初始化"
- )
- return
+ auth_header = {"Authorization": "Bearer; {}".format(self.access_token)}
+ async with websockets.connect(
+ self.ws_url, additional_headers=auth_header
+ ) as websocket:
+ # Prepare request data
+ request_params = self._construct_request(str(uuid.uuid4()))
+ payload_bytes = str.encode(json.dumps(request_params))
+ payload_bytes = gzip.compress(payload_bytes)
+ full_client_request = self._generate_header()
+ full_client_request.extend(
+ (len(payload_bytes)).to_bytes(4, "big")
+ ) # payload size(4 bytes)
+ full_client_request.extend(payload_bytes) # payload
- try:
- logger.bind(tag=TAG).info("开始启动ASR监听...")
- asyncio.run_coroutine_threadsafe(
- self._forward_asr_results(), loop=self.conn.loop
- )
- logger.bind(tag=TAG).info("ASR监听已启动")
- except Exception as e:
- logger.bind(tag=TAG).error(f"启动ASR监听线程失败: {e}")
- except Exception as e:
- logger.bind(tag=TAG).error(f"ASR监听线程发生未预期的错误: {e}")
+ # Send header and metadata
+ # full_client_request
+ await websocket.send(full_client_request)
+ res = await websocket.recv()
+ result = parse_response(res)
+ if (
+ "payload_msg" in result
+ and result["payload_msg"]["code"] != self.success_code
+ ):
+ logger.bind(tag=TAG).error(f"ASR error: {result}")
+ return None
- async def _forward_asr_results(self):
- try:
- while not self.conn.stop_event.is_set():
- try:
- if self.asr_ws is None:
- # 检查是否需要重连
- async with self.reconnect_lock:
- current_time = asyncio.get_event_loop().time()
- if (
- current_time - self.last_reconnect_time
- < self.reconnect_cooldown
- ):
- await asyncio.sleep(1)
- continue
-
- if self.reconnect_count >= self.max_reconnect_count:
- logger.bind(tag=TAG).error(
- "达到最大重连次数限制,停止重连"
- )
- await asyncio.sleep(self.reconnect_cooldown)
- self.reconnect_count = 0
- continue
-
- self.last_reconnect_time = current_time
- self.reconnect_count += 1
- logger.bind(tag=TAG).info(
- f"尝试重新连接ASR服务... (第{self.reconnect_count}次)"
- )
- await self.open_audio_channels(self.conn)
- continue
-
- # 使用锁来确保同一时间只有一个协程在接收数据
- async with self.recv_lock:
- response = await self.asr_ws.recv()
- result = self.parse_response(response)
-
- # 检查是否需要重连
- if result.get("need_reconnect", False):
- logger.bind(tag=TAG).info(
- "检测到需要重连的错误,准备重新连接..."
+ for seq, (chunk, last) in enumerate(
+ self.slice_data(audio_data, segment_size), 1
+ ):
+ if last:
+ audio_only_request = self._generate_header(
+ message_type=CLIENT_AUDIO_ONLY_REQUEST,
+ message_type_specific_flags=NEG_SEQUENCE,
)
- if self.asr_ws is not None:
- try:
- await self.asr_ws.close()
- except Exception as e:
- logger.bind(tag=TAG).warning(
- f"关闭旧连接时发生错误: {e}"
- )
- finally:
- self.asr_ws = None
- continue
+ else:
+ audio_only_request = self._generate_header(
+ message_type=CLIENT_AUDIO_ONLY_REQUEST
+ )
+ payload_bytes = gzip.compress(chunk)
+ audio_only_request.extend(
+ (len(payload_bytes)).to_bytes(4, "big")
+ ) # payload size(4 bytes)
+ audio_only_request.extend(payload_bytes) # payload
+ # Send audio data
+ await websocket.send(audio_only_request)
- if "payload_msg" in result:
- if "result" in result["payload_msg"]:
- # 检查是否有utterances并且definite为True
- utterances = result["payload_msg"]["result"][0].get(
- "utterances", []
- )
- for utterance in utterances:
- if utterance.get("definite", False):
- self.text = utterance["text"]
- await self.handle_voice_stop(None)
- break
+ # Receive response
+ response = await websocket.recv()
+ result = parse_response(response)
- except websockets.ConnectionClosed:
- logger.bind(tag=TAG).debug("ASR服务连接已关闭,准备重连...")
- # 确保关闭旧连接
- if self.asr_ws is not None:
- try:
- await self.asr_ws.close()
- except Exception as e:
- logger.bind(tag=TAG).warning(f"关闭旧连接时发生错误: {e}")
- finally:
- self.asr_ws = None
-
- # 等待冷却时间
- await asyncio.sleep(self.reconnect_cooldown)
- continue
-
- except Exception as e:
- if not self.conn.stop_event.is_set():
- logger.bind(tag=TAG).error(f"ASR监听发生错误: {e}")
- await asyncio.sleep(self.retry_delay)
- continue
+ if (
+ "payload_msg" in result
+ and result["payload_msg"]["code"] == self.success_code
+ ):
+ if len(result["payload_msg"]["result"]) > 0:
+ return result["payload_msg"]["result"][0]["text"]
+ return None
+ else:
+ logger.bind(tag=TAG).error(f"ASR error: {result}")
+ return None
except Exception as e:
- logger.bind(tag=TAG).error(f"ASR监听线程发生错误: {e}")
- # 确保在发生严重错误时也能继续尝试重连
- if not self.conn.stop_event.is_set():
- await asyncio.sleep(self.retry_delay)
- await self._forward_asr_results() # 递归重试
+ logger.bind(tag=TAG).error(f"ASR request failed: {e}", exc_info=True)
+ return None
- async def speech_to_text(self, opus_data, session_id):
- result = self.text
- self.text = "" # 清空text
- return result, None
-
- def parse_response(self, res: bytes) -> dict:
+ @staticmethod
+ def slice_data(data: bytes, chunk_size: int) -> (list, bool):
"""
- 解析 ASR 服务返回的二进制响应。
- 根据协议格式解析头部和 payload,若采用 GZIP 压缩则先解压,再根据 JSON 反序列化。
+ slice data
+ :param data: wav data
+ :param chunk_size: the segment size in one request
+ :return: segment data, last flag
"""
- protocol_version = res[0] >> 4
- header_size = res[0] & 0x0F
- message_type = res[1] >> 4
- serialization_method = res[2] >> 4
- message_compression = res[2] & 0x0F
- payload = res[header_size * 4 :]
- result = {}
- payload_msg = None
- payload_size = 0
-
- if message_type == SERVER_FULL_RESPONSE:
- payload_size = int.from_bytes(payload[:4], "big", signed=True)
- payload_msg = payload[4:]
- elif message_type == SERVER_ACK:
- seq = int.from_bytes(payload[:4], "big", signed=True)
- result["seq"] = seq
- if len(payload) >= 8:
- payload_size = int.from_bytes(payload[4:8], "big", signed=False)
- payload_msg = payload[8:]
- elif message_type == SERVER_ERROR_RESPONSE:
- code = int.from_bytes(payload[:4], "big", signed=False)
- result["code"] = code
- payload_size = int.from_bytes(payload[4:8], "big", signed=False)
- payload_msg = payload[8:]
-
- if payload_msg is None:
- return result
- if message_compression == GZIP_COMPRESSION:
- payload_msg = gzip.decompress(payload_msg)
- if serialization_method == JSON_SERIALIZATION:
- payload_msg = json.loads(payload_msg.decode("utf-8"))
+ data_len = len(data)
+ offset = 0
+ while offset + chunk_size < data_len:
+ yield data[offset : offset + chunk_size], False
+ offset += chunk_size
else:
- payload_msg = payload_msg.decode("utf-8")
- result["payload_msg"] = payload_msg
- result["payload_size"] = payload_size
+ yield data[offset:data_len], True
- # 错误码处理
- if "code" in result:
- error_code = result["code"]
- error_message = ""
+ async def speech_to_text(
+ self, opus_data: List[bytes], session_id: str
+ ) -> Tuple[Optional[str], Optional[str]]:
+ """将语音数据转换为文本"""
- if error_code == 1000:
- error_message = "成功"
- elif error_code == 1001:
- error_message = "请求参数无效:请求参数缺失必需字段/字段值无效/重复请求"
- elif error_code == 1002:
- error_message = "无访问权限:token无效/过期/无权访问指定服务"
- elif error_code == 1003:
- error_message = "访问超频:当前appid访问QPS超出设定阈值"
- elif error_code == 1004:
- error_message = "访问超额:当前appid访问次数超出限制"
- elif error_code == 1005:
- error_message = "服务器繁忙:服务过载,无法处理当前请求"
- elif error_code == 1010:
- error_message = "音频过长:音频数据时长超出阈值"
- elif error_code == 1011:
- error_message = "音频过大:音频数据大小超出阈值"
- elif error_code == 1012:
- error_message = "音频格式无效:音频header有误/无法进行音频解码"
- elif error_code == 1013:
- error_message = "音频静音:音频未识别出任何文本结果"
- elif error_code >= 1020 and error_code <= 1022:
- error_message = "识别相关错误:需要重连"
- if error_code == 1020:
- error_message = "识别等待超时:等待下一包就绪超时"
- elif error_code == 1021:
- error_message = "识别处理超时:识别处理过程超时"
- elif error_code == 1022:
- error_message = "识别错误:识别过程中发生错误"
+ file_path = None
+ try:
+ # 合并所有opus数据包
+ if self.audio_format == "pcm":
+ pcm_data = opus_data
else:
- error_message = "未知错误:未归类错误"
+ pcm_data = self.decode_opus(opus_data)
+ combined_pcm_data = b"".join(pcm_data)
- logger.bind(tag=TAG).debug(
- f"ASR错误: {error_message} (错误码: {error_code})"
- )
+ # 判断是否保存为WAV文件
+ if self.delete_audio_file:
+ pass
+ else:
+ file_path = self.save_audio_to_file(pcm_data, session_id)
- # 如果是识别相关错误,标记需要重连
- if error_code >= 1020 or error_code == 1001:
- result["need_reconnect"] = True
+ # 直接使用PCM数据
+ # 计算分段大小 (单声道, 16bit, 16kHz采样率)
+ size_per_sec = 1 * 2 * 16000 # nchannels * sampwidth * framerate
+ segment_size = int(size_per_sec * self.seg_duration / 1000)
- return result
-
- async def close_session(self):
- """关闭当前会话"""
- async with self._session_lock:
- if not self._session_started:
- logger.bind(tag=TAG).warning("尝试关闭未开始的会话")
- return
-
- if self._session_finished:
- logger.bind(tag=TAG).warning(
- f"会话 {self._current_session_id} 已经关闭"
+ # 语音识别
+ start_time = time.time()
+ text = await self._send_request(combined_pcm_data, segment_size)
+ if text:
+ logger.bind(tag=TAG).debug(
+ f"语音识别耗时: {time.time() - start_time:.3f}s | 结果: {text}"
)
- return
+ return text, file_path
+ return "", file_path
- try:
- if self.asr_ws is not None:
- await self.asr_ws.close()
- except Exception as e:
- logger.bind(tag=TAG).warning(f"关闭WebSocket连接时发生错误: {e}")
- finally:
- self.asr_ws = None
- self._session_finished = True
- self._session_started = False
- self._current_session_id = None
- # 重置重连计数
- self.reconnect_count = 0
-
- async def close(self):
- """资源清理方法"""
- await self.close_session()
+ except Exception as e:
+ logger.bind(tag=TAG).error(f"语音识别失败: {e}", exc_info=True)
+ return "", file_path
diff --git a/main/xiaozhi-server/core/providers/asr/doubao_stream.py b/main/xiaozhi-server/core/providers/asr/doubao_stream.py
new file mode 100644
index 00000000..da4098ad
--- /dev/null
+++ b/main/xiaozhi-server/core/providers/asr/doubao_stream.py
@@ -0,0 +1,534 @@
+import json
+import gzip
+import uuid
+import asyncio
+import websockets
+import opuslib_next
+from core.providers.asr.base import ASRProviderBase
+from config.logger import setup_logging
+from core.providers.asr.dto.dto import InterfaceType
+import threading
+
+TAG = __name__
+logger = setup_logging()
+
+CLIENT_FULL_REQUEST = 0b0001
+CLIENT_AUDIO_ONLY_REQUEST = 0b0010
+SERVER_FULL_RESPONSE = 0b1001
+SERVER_ACK = 0b1011
+SERVER_ERROR_RESPONSE = 0b1111
+NO_SEQUENCE = 0b0000
+NEG_SEQUENCE = 0b0010
+JSON_SERIALIZATION = 0b0001
+GZIP_COMPRESSION = 0b0001
+PROTOCOL_VERSION = 0b0001
+
+
+class ASRProvider(ASRProviderBase):
+ def __init__(self, config, delete_audio_file):
+ super().__init__()
+ self.interface_type = InterfaceType.STREAM
+ self.config = config
+ self.text = ""
+ self.max_retries = 3
+ self.retry_delay = 2 # 重试延迟秒数
+ self.recv_lock = asyncio.Lock() # 添加接收锁
+ self.reconnect_lock = asyncio.Lock() # 添加重连锁
+ self.last_reconnect_time = 0 # 上次重连时间
+ self.reconnect_cooldown = 1 # 增加重连冷却时间到10秒
+ self.reconnect_count = 0 # 当前重连次数
+ self.max_reconnect_count = 3 # 减少最大重连次数到3次
+ self.asr_thread = None # ASR监听线程
+ self.thread_lock = threading.Lock() # 线程管理锁
+ self.is_reconnecting = False # 添加重连状态标志
+
+ # 添加会话管理相关属性
+ self._session_lock = asyncio.Lock() # 会话操作的并发锁
+ self._current_session_id = None # 当前会话ID
+ self._session_started = False # 会话是否已开始
+ self._session_finished = False # 会话是否已结束
+ self._session_close_event = asyncio.Event() # 添加会话关闭事件
+
+ self.appid = str(config.get("appid"))
+ self.cluster = config.get("cluster")
+ self.access_token = config.get("access_token")
+ self.boosting_table_name = config.get("boosting_table_name", "")
+ self.correct_table_name = config.get("correct_table_name", "")
+ self.output_dir = config.get("output_dir", "temp/")
+ self.delete_audio_file = delete_audio_file
+
+ self.ws_url = "wss://openspeech.bytedance.com/api/v2/asr"
+ self.uid = config.get("uid", "streaming_asr_service")
+ self.workflow = config.get(
+ "workflow", "audio_in,resample,partition,vad,fe,decode,itn,nlu_punctuate"
+ )
+ self.result_type = config.get("result_type", "single")
+ self.format = config.get("format", "raw")
+ self.codec = config.get("codec", "pcm")
+ self.rate = config.get("sample_rate", 16000)
+ self.language = config.get("language", "zh-CN")
+ self.bits = config.get("bits", 16)
+ self.channel = config.get("channel", 1)
+ self.auth_method = config.get("auth_method", "token")
+ self.secret = config.get("secret", "access_secret")
+ self.decoder = opuslib_next.Decoder(16000, 1)
+ self.asr_ws = None
+ self.forward_task = None
+ self.conn = None
+
+ ###################################################################################
+ # 豆包流式ASR重写父类的方法--开始
+ ###################################################################################
+ async def open_audio_channels(self, conn):
+ await super().open_audio_channels(conn)
+
+ async with self._session_lock:
+ # 如果正在重连,等待重连完成
+ if self.is_reconnecting:
+ logger.bind(tag=TAG).info("等待当前重连完成...")
+ await self._session_close_event.wait()
+ self._session_close_event.clear()
+
+ # 如果已有会话未结束,先关闭它
+ if self._session_started and not self._session_finished:
+ logger.bind(tag=TAG).warning(
+ f"发现未关闭的会话 {self._current_session_id},正在关闭..."
+ )
+ if self.asr_ws is not None:
+ try:
+ await self.asr_ws.close()
+ except Exception as e:
+ logger.bind(tag=TAG).warning(f"关闭旧连接时发生错误: {e}")
+ finally:
+ self.asr_ws = None
+ self._session_finished = True
+ self._session_close_event.set()
+
+ # 重置会话状态
+ self._current_session_id = str(uuid.uuid4())
+ self._session_started = True
+ self._session_finished = False
+ self.is_reconnecting = True
+
+ try:
+ retry_count = 0
+ while retry_count < self.max_retries:
+ try:
+ headers = (
+ self.token_auth() if self.auth_method == "token" else None
+ )
+ self.asr_ws = await websockets.connect(
+ self.ws_url,
+ additional_headers=headers,
+ max_size=1000000000,
+ ping_interval=None,
+ ping_timeout=None,
+ close_timeout=10,
+ )
+
+ # 发送初始化请求
+ request_params = self.construct_request(
+ self._current_session_id
+ )
+ try:
+ payload_bytes = str.encode(json.dumps(request_params))
+ payload_bytes = gzip.compress(payload_bytes)
+ full_client_request = self.generate_header()
+ full_client_request.extend(
+ (len(payload_bytes)).to_bytes(4, "big")
+ )
+ full_client_request.extend(payload_bytes)
+ await self.asr_ws.send(full_client_request)
+ except Exception as e:
+ logger.bind(tag=TAG).error(f"发送初始化请求失败: {e}")
+ raise e
+
+ # 等待初始化响应
+ try:
+ init_res = await self.asr_ws.recv()
+ self.parse_response(init_res)
+ except Exception as e:
+ logger.bind(tag=TAG).error(f"ASR服务初始化失败: {e}")
+ raise e
+
+ # 启动接收ASR结果的异步任务
+ with self.thread_lock:
+ if (
+ self.asr_thread is None
+ or not self.asr_thread.is_alive()
+ ):
+ logger.bind(tag=TAG).info("创建新的ASR监听线程...")
+ self.asr_thread = threading.Thread(
+ target=self._start_monitor_asr_response_thread,
+ daemon=True,
+ )
+ self.asr_thread.start()
+ # 等待一小段时间确保线程启动
+ await asyncio.sleep(0.1)
+ if not self.asr_thread.is_alive():
+ logger.bind(tag=TAG).error("ASR监听线程启动失败")
+ raise Exception("ASR监听线程启动失败")
+ logger.bind(tag=TAG).info("ASR监听线程已启动")
+ return
+
+ except websockets.exceptions.WebSocketException as e:
+ retry_count += 1
+ if retry_count < self.max_retries:
+ logger.bind(tag=TAG).warning(
+ f"WebSocket连接失败,正在进行第{retry_count}次重试: {e}"
+ )
+ await asyncio.sleep(self.retry_delay)
+ else:
+ logger.bind(tag=TAG).warning(
+ f"WebSocket连接失败,已达到最大重试次数: {e}"
+ )
+ raise
+ except Exception as e:
+ logger.bind(tag=TAG).error(f"WebSocket连接发生未知错误: {e}")
+ raise
+ finally:
+ self.is_reconnecting = False
+ self._session_close_event.set()
+
+ async def receive_audio(self, audio, _):
+ if not isinstance(audio, bytes):
+ return
+
+ try:
+ # 解码opus得到PCM数据
+ pcm_frame = self.decoder.decode(audio, 960)
+ payload = gzip.compress(pcm_frame)
+ audio_request = bytearray(self.generate_audio_default_header())
+ audio_request.extend(len(payload).to_bytes(4, "big"))
+ audio_request.extend(payload)
+ if self.asr_ws:
+ await self.asr_ws.send(audio_request)
+ except Exception as e:
+ logger.bind(tag=TAG).debug(f"发送音频数据时发生错误: {e}")
+
+ ###################################################################################
+ # 豆包流式ASR重写父类的方法--结束
+ ###################################################################################
+
+ def construct_request(self, reqid):
+ req = {
+ "app": {
+ "appid": self.appid,
+ "cluster": self.cluster,
+ "token": self.access_token,
+ },
+ "user": {"uid": self.uid},
+ "request": {
+ "reqid": reqid,
+ "workflow": self.workflow,
+ "show_utterances": True,
+ "result_type": self.result_type,
+ "sequence": 1,
+ "boosting_table_name": self.boosting_table_name,
+ "correct_table_name": self.correct_table_name,
+ },
+ "audio": {
+ "format": self.format,
+ "codec": self.codec,
+ "rate": self.rate,
+ "language": self.language,
+ "bits": self.bits,
+ "channel": self.channel,
+ },
+ }
+ return req
+
+ def token_auth(self):
+ return {"Authorization": f"Bearer; {self.access_token}"}
+
+ def generate_header(
+ self,
+ version=PROTOCOL_VERSION,
+ message_type=CLIENT_FULL_REQUEST,
+ message_type_specific_flags=NO_SEQUENCE,
+ serial_method=JSON_SERIALIZATION,
+ compression_type=GZIP_COMPRESSION,
+ reserved_data=0x00,
+ extension_header: bytes = b"",
+ ):
+ """
+ 生成协议头:
+ - 第1字节:高4位:协议版本,低4位:头部大小(单位 4 字节)
+ - 第2字节:高4位:消息类型,低4位:消息类型特定标志
+ - 第3字节:高4位:序列化方式,低4位:压缩方式
+ - 第4字节:保留字段
+ - 后续:扩展头(如果有)
+ """
+ header = bytearray()
+ 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_full_default_header(self):
+ # full client request 默认头
+ return self.generate_header(
+ version=PROTOCOL_VERSION,
+ message_type=CLIENT_FULL_REQUEST,
+ message_type_specific_flags=NO_SEQUENCE,
+ serial_method=JSON_SERIALIZATION,
+ compression_type=GZIP_COMPRESSION,
+ )
+
+ def generate_audio_default_header(self):
+ # 普通音频片段请求
+ return self.generate_header(
+ version=PROTOCOL_VERSION,
+ message_type=CLIENT_AUDIO_ONLY_REQUEST,
+ message_type_specific_flags=NO_SEQUENCE,
+ serial_method=JSON_SERIALIZATION,
+ compression_type=GZIP_COMPRESSION,
+ )
+
+ def generate_last_audio_default_header(self):
+ # 最后一个音频片段标志
+ return self.generate_header(
+ version=PROTOCOL_VERSION,
+ message_type=CLIENT_AUDIO_ONLY_REQUEST,
+ message_type_specific_flags=NEG_SEQUENCE, # 用 NEG_SEQUENCE 表示结束
+ serial_method=JSON_SERIALIZATION,
+ compression_type=GZIP_COMPRESSION,
+ )
+
+ def _start_monitor_asr_response_thread(self):
+ # 初始化链接
+ try:
+ with self.thread_lock:
+ if self.conn is None or self.conn.loop is None:
+ logger.bind(tag=TAG).error(
+ "无法启动ASR监听线程:conn或loop未初始化"
+ )
+ return
+
+ try:
+ logger.bind(tag=TAG).info("开始启动ASR监听...")
+ asyncio.run_coroutine_threadsafe(
+ self._forward_asr_results(), loop=self.conn.loop
+ )
+ logger.bind(tag=TAG).info("ASR监听已启动")
+ except Exception as e:
+ logger.bind(tag=TAG).error(f"启动ASR监听线程失败: {e}")
+ except Exception as e:
+ logger.bind(tag=TAG).error(f"ASR监听线程发生未预期的错误: {e}")
+
+ async def _forward_asr_results(self):
+ try:
+ while not self.conn.stop_event.is_set():
+ try:
+ if self.asr_ws is None:
+ # 检查是否需要重连
+ async with self.reconnect_lock:
+ current_time = asyncio.get_event_loop().time()
+ if (
+ current_time - self.last_reconnect_time
+ < self.reconnect_cooldown
+ ):
+ await asyncio.sleep(1)
+ continue
+
+ if self.reconnect_count >= self.max_reconnect_count:
+ logger.bind(tag=TAG).error(
+ "达到最大重连次数限制,停止重连"
+ )
+ await asyncio.sleep(self.reconnect_cooldown)
+ self.reconnect_count = 0
+ continue
+
+ self.last_reconnect_time = current_time
+ self.reconnect_count += 1
+ logger.bind(tag=TAG).info(
+ f"尝试重新连接ASR服务... (第{self.reconnect_count}次)"
+ )
+ await self.open_audio_channels(self.conn)
+ continue
+
+ # 使用锁来确保同一时间只有一个协程在接收数据
+ async with self.recv_lock:
+ response = await self.asr_ws.recv()
+ result = self.parse_response(response)
+
+ # 检查是否需要重连
+ if result.get("need_reconnect", False):
+ logger.bind(tag=TAG).info(
+ "检测到需要重连的错误,准备重新连接..."
+ )
+ if self.asr_ws is not None:
+ try:
+ await self.asr_ws.close()
+ except Exception as e:
+ logger.bind(tag=TAG).warning(
+ f"关闭旧连接时发生错误: {e}"
+ )
+ finally:
+ self.asr_ws = None
+ continue
+
+ if "payload_msg" in result:
+ if "result" in result["payload_msg"]:
+ # 检查是否有utterances并且definite为True
+ utterances = result["payload_msg"]["result"][0].get(
+ "utterances", []
+ )
+ for utterance in utterances:
+ if utterance.get("definite", False):
+ self.text = utterance["text"]
+ await self.handle_voice_stop(None)
+ break
+
+ except websockets.ConnectionClosed:
+ logger.bind(tag=TAG).debug("ASR服务连接已关闭,准备重连...")
+ # 确保关闭旧连接
+ if self.asr_ws is not None:
+ try:
+ await self.asr_ws.close()
+ except Exception as e:
+ logger.bind(tag=TAG).warning(f"关闭旧连接时发生错误: {e}")
+ finally:
+ self.asr_ws = None
+
+ # 等待冷却时间
+ await asyncio.sleep(self.reconnect_cooldown)
+ continue
+
+ except Exception as e:
+ if not self.conn.stop_event.is_set():
+ logger.bind(tag=TAG).error(f"ASR监听发生错误: {e}")
+ await asyncio.sleep(self.retry_delay)
+ continue
+
+ except Exception as e:
+ logger.bind(tag=TAG).error(f"ASR监听线程发生错误: {e}")
+ # 确保在发生严重错误时也能继续尝试重连
+ if not self.conn.stop_event.is_set():
+ await asyncio.sleep(self.retry_delay)
+ await self._forward_asr_results() # 递归重试
+
+ async def speech_to_text(self, opus_data, session_id):
+ result = self.text
+ self.text = "" # 清空text
+ return result, None
+
+ def parse_response(self, res: bytes) -> dict:
+ """
+ 解析 ASR 服务返回的二进制响应。
+ 根据协议格式解析头部和 payload,若采用 GZIP 压缩则先解压,再根据 JSON 反序列化。
+ """
+ protocol_version = res[0] >> 4
+ header_size = res[0] & 0x0F
+ message_type = res[1] >> 4
+ serialization_method = res[2] >> 4
+ message_compression = res[2] & 0x0F
+ payload = res[header_size * 4 :]
+ result = {}
+ payload_msg = None
+ payload_size = 0
+
+ if message_type == SERVER_FULL_RESPONSE:
+ payload_size = int.from_bytes(payload[:4], "big", signed=True)
+ payload_msg = payload[4:]
+ elif message_type == SERVER_ACK:
+ seq = int.from_bytes(payload[:4], "big", signed=True)
+ result["seq"] = seq
+ if len(payload) >= 8:
+ payload_size = int.from_bytes(payload[4:8], "big", signed=False)
+ payload_msg = payload[8:]
+ elif message_type == SERVER_ERROR_RESPONSE:
+ code = int.from_bytes(payload[:4], "big", signed=False)
+ result["code"] = code
+ payload_size = int.from_bytes(payload[4:8], "big", signed=False)
+ payload_msg = payload[8:]
+
+ if payload_msg is None:
+ return result
+ if message_compression == GZIP_COMPRESSION:
+ payload_msg = gzip.decompress(payload_msg)
+ if serialization_method == JSON_SERIALIZATION:
+ payload_msg = json.loads(payload_msg.decode("utf-8"))
+ else:
+ payload_msg = payload_msg.decode("utf-8")
+ result["payload_msg"] = payload_msg
+ result["payload_size"] = payload_size
+
+ # 错误码处理
+ if "code" in result:
+ error_code = result["code"]
+ error_message = ""
+
+ if error_code == 1000:
+ error_message = "成功"
+ elif error_code == 1001:
+ error_message = "请求参数无效:请求参数缺失必需字段/字段值无效/重复请求"
+ elif error_code == 1002:
+ error_message = "无访问权限:token无效/过期/无权访问指定服务"
+ elif error_code == 1003:
+ error_message = "访问超频:当前appid访问QPS超出设定阈值"
+ elif error_code == 1004:
+ error_message = "访问超额:当前appid访问次数超出限制"
+ elif error_code == 1005:
+ error_message = "服务器繁忙:服务过载,无法处理当前请求"
+ elif error_code == 1010:
+ error_message = "音频过长:音频数据时长超出阈值"
+ elif error_code == 1011:
+ error_message = "音频过大:音频数据大小超出阈值"
+ elif error_code == 1012:
+ error_message = "音频格式无效:音频header有误/无法进行音频解码"
+ elif error_code == 1013:
+ error_message = "音频静音:音频未识别出任何文本结果"
+ elif error_code >= 1020 and error_code <= 1022:
+ error_message = "识别相关错误:需要重连"
+ if error_code == 1020:
+ error_message = "识别等待超时:等待下一包就绪超时"
+ elif error_code == 1021:
+ error_message = "识别处理超时:识别处理过程超时"
+ elif error_code == 1022:
+ error_message = "识别错误:识别过程中发生错误"
+ else:
+ error_message = "未知错误:未归类错误"
+
+ logger.bind(tag=TAG).debug(
+ f"ASR错误: {error_message} (错误码: {error_code})"
+ )
+
+ # 如果是识别相关错误,标记需要重连
+ if error_code >= 1020 or error_code == 1001:
+ result["need_reconnect"] = True
+
+ return result
+
+ async def close_session(self):
+ """关闭当前会话"""
+ async with self._session_lock:
+ if not self._session_started:
+ logger.bind(tag=TAG).warning("尝试关闭未开始的会话")
+ return
+
+ if self._session_finished:
+ logger.bind(tag=TAG).warning(
+ f"会话 {self._current_session_id} 已经关闭"
+ )
+ return
+
+ try:
+ if self.asr_ws is not None:
+ await self.asr_ws.close()
+ except Exception as e:
+ logger.bind(tag=TAG).warning(f"关闭WebSocket连接时发生错误: {e}")
+ finally:
+ self.asr_ws = None
+ self._session_finished = True
+ self._session_started = False
+ self._current_session_id = None
+ # 重置重连计数
+ self.reconnect_count = 0
+
+ async def close(self):
+ """资源清理方法"""
+ await self.close_session()
diff --git a/main/xiaozhi-server/core/providers/tts/base.py b/main/xiaozhi-server/core/providers/tts/base.py
index 9c094f24..70c68ee7 100644
--- a/main/xiaozhi-server/core/providers/tts/base.py
+++ b/main/xiaozhi-server/core/providers/tts/base.py
@@ -52,6 +52,7 @@ class TTSProviderBase(ABC):
)
self.first_sentence_punctuations = (
",",
+ "~",
"~",
"、",
",",