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 1/2] =?UTF-8?q?update:=E5=8C=BA=E5=88=86=E8=B1=86=E5=8C=85?= =?UTF-8?q?ASR=E6=8C=89=E6=AC=A1=E6=94=B6=E8=B4=B9=E5=92=8C=E6=8C=89?= =?UTF-8?q?=E6=97=B6=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 = ( ",", + "~", "~", "、", ",", From 109811199d89fb9d33258c9ebcffbde472446aac Mon Sep 17 00:00:00 2001 From: hrz <1710360675@qq.com> Date: Tue, 3 Jun 2025 17:32:18 +0800 Subject: [PATCH 2/2] =?UTF-8?q?=E6=9B=B4=E6=96=B0=EF=BC=9A=E6=99=BA?= =?UTF-8?q?=E6=8E=A7=E5=8F=B0=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 --- .../main/resources/db/changelog/db.changelog-master.yaml | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) diff --git a/main/manager-api/src/main/resources/db/changelog/db.changelog-master.yaml b/main/manager-api/src/main/resources/db/changelog/db.changelog-master.yaml index 437e8ea6..1bd070aa 100755 --- a/main/manager-api/src/main/resources/db/changelog/db.changelog-master.yaml +++ b/main/manager-api/src/main/resources/db/changelog/db.changelog-master.yaml @@ -176,4 +176,11 @@ databaseChangeLog: changes: - sqlFile: encoding: utf8 - path: classpath:db/changelog/202506010920.sql \ No newline at end of file + path: classpath:db/changelog/202506010920.sql + - changeSet: + id: 202506031639 + author: hrz + changes: + - sqlFile: + encoding: utf8 + path: classpath:db/changelog/202506031639.sql \ No newline at end of file