mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-26 17:13:54 +08:00
Merge pull request #1459 from xinnan-tech/fix-doubao-asr
update:区分豆包ASR按次收费和按时收费接口
This commit is contained in:
@@ -227,7 +227,7 @@ public interface Constant {
|
|||||||
/**
|
/**
|
||||||
* 版本号
|
* 版本号
|
||||||
*/
|
*/
|
||||||
public static final String VERSION = "0.5.2";
|
public static final String VERSION = "0.5.4";
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* 无效固件URL
|
* 无效固件URL
|
||||||
|
|||||||
@@ -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';
|
||||||
@@ -177,3 +177,10 @@ databaseChangeLog:
|
|||||||
- sqlFile:
|
- sqlFile:
|
||||||
encoding: utf8
|
encoding: utf8
|
||||||
path: classpath:db/changelog/202506010920.sql
|
path: classpath:db/changelog/202506010920.sql
|
||||||
|
- changeSet:
|
||||||
|
id: 202506031639
|
||||||
|
author: hrz
|
||||||
|
changes:
|
||||||
|
- sqlFile:
|
||||||
|
encoding: utf8
|
||||||
|
path: classpath:db/changelog/202506031639.sql
|
||||||
@@ -31,7 +31,7 @@
|
|||||||
<span class="menu-text">大语言模型</span>
|
<span class="menu-text">大语言模型</span>
|
||||||
</el-menu-item>
|
</el-menu-item>
|
||||||
<el-menu-item index="vllm">
|
<el-menu-item index="vllm">
|
||||||
<span class="menu-text">视觉大语言模型</span>
|
<span class="menu-text">视觉大模型</span>
|
||||||
</el-menu-item>
|
</el-menu-item>
|
||||||
<el-menu-item index="intent">
|
<el-menu-item index="intent">
|
||||||
<span class="menu-text">意图识别</span>
|
<span class="menu-text">意图识别</span>
|
||||||
@@ -176,7 +176,7 @@ export default {
|
|||||||
vad: '语言活动检测模型(VAD)',
|
vad: '语言活动检测模型(VAD)',
|
||||||
asr: '语音识别模型(ASR)',
|
asr: '语音识别模型(ASR)',
|
||||||
llm: '大语言模型(LLM)',
|
llm: '大语言模型(LLM)',
|
||||||
vllm: '视觉大语言模型(VLLM)',
|
vllm: '视觉大模型(VLLM)',
|
||||||
intent: '意图识别模型(Intent)',
|
intent: '意图识别模型(Intent)',
|
||||||
tts: '语音合成模型(TTS)',
|
tts: '语音合成模型(TTS)',
|
||||||
memory: '记忆模型(Memory)'
|
memory: '记忆模型(Memory)'
|
||||||
|
|||||||
@@ -177,7 +177,7 @@ export default {
|
|||||||
{ label: '语音活动检测(VAD)', key: 'vadModelId', type: 'VAD' },
|
{ label: '语音活动检测(VAD)', key: 'vadModelId', type: 'VAD' },
|
||||||
{ label: '语音识别(ASR)', key: 'asrModelId', type: 'ASR' },
|
{ label: '语音识别(ASR)', key: 'asrModelId', type: 'ASR' },
|
||||||
{ label: '大语言模型(LLM)', key: 'llmModelId', type: 'LLM' },
|
{ 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: '意图识别(Intent)', key: 'intentModelId', type: 'Intent' },
|
||||||
{ label: '记忆(Memory)', key: 'memModelId', type: 'Memory' },
|
{ label: '记忆(Memory)', key: 'memModelId', type: 'Memory' },
|
||||||
{ label: '语音合成(TTS)', key: 'ttsModelId', type: 'TTS' },
|
{ label: '语音合成(TTS)', key: 'ttsModelId', type: 'TTS' },
|
||||||
|
|||||||
@@ -264,6 +264,8 @@ ASR:
|
|||||||
DoubaoASR:
|
DoubaoASR:
|
||||||
# 可以在这里申请相关Key等信息
|
# 可以在这里申请相关Key等信息
|
||||||
# https://console.volcengine.com/speech/app
|
# https://console.volcengine.com/speech/app
|
||||||
|
# DoubaoASR和DoubaoStreamASR的区别是:DoubaoASR是按次收费,DoubaoStreamASR是按时收费
|
||||||
|
# 一般来说按次收费的更便宜,但是DoubaoStreamASR使用了大模型技术,效果更好
|
||||||
type: doubao
|
type: doubao
|
||||||
appid: 你的火山引擎语音合成服务appid
|
appid: 你的火山引擎语音合成服务appid
|
||||||
access_token: 你的火山引擎语音合成服务access_token
|
access_token: 你的火山引擎语音合成服务access_token
|
||||||
@@ -272,6 +274,19 @@ ASR:
|
|||||||
boosting_table_name: (选填)你的热词文件名称
|
boosting_table_name: (选填)你的热词文件名称
|
||||||
correct_table_name: (选填)你的替换词文件名称
|
correct_table_name: (选填)你的替换词文件名称
|
||||||
output_dir: tmp/
|
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:
|
TencentASR:
|
||||||
# token申请地址:https://console.cloud.tencent.com/cam/capi
|
# token申请地址:https://console.cloud.tencent.com/cam/capi
|
||||||
# 免费领取资源:https://console.cloud.tencent.com/asr/resourcebundle
|
# 免费领取资源:https://console.cloud.tencent.com/asr/resourcebundle
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ from loguru import logger
|
|||||||
from config.config_loader import load_config
|
from config.config_loader import load_config
|
||||||
from config.settings import check_config_file
|
from config.settings import check_config_file
|
||||||
|
|
||||||
SERVER_VERSION = "0.5.2"
|
SERVER_VERSION = "0.5.4"
|
||||||
_logger_initialized = False
|
_logger_initialized = False
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,534 +1,267 @@
|
|||||||
|
import time
|
||||||
|
import os
|
||||||
|
import uuid
|
||||||
import json
|
import json
|
||||||
import gzip
|
import gzip
|
||||||
import uuid
|
|
||||||
import asyncio
|
|
||||||
import websockets
|
import websockets
|
||||||
import opuslib_next
|
|
||||||
from core.providers.asr.base import ASRProviderBase
|
|
||||||
from config.logger import setup_logging
|
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
|
from core.providers.asr.dto.dto import InterfaceType
|
||||||
import threading
|
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
|
|
||||||
CLIENT_FULL_REQUEST = 0b0001
|
CLIENT_FULL_REQUEST = 0b0001
|
||||||
CLIENT_AUDIO_ONLY_REQUEST = 0b0010
|
CLIENT_AUDIO_ONLY_REQUEST = 0b0010
|
||||||
|
|
||||||
|
NO_SEQUENCE = 0b0000
|
||||||
|
NEG_SEQUENCE = 0b0010
|
||||||
|
|
||||||
SERVER_FULL_RESPONSE = 0b1001
|
SERVER_FULL_RESPONSE = 0b1001
|
||||||
SERVER_ACK = 0b1011
|
SERVER_ACK = 0b1011
|
||||||
SERVER_ERROR_RESPONSE = 0b1111
|
SERVER_ERROR_RESPONSE = 0b1111
|
||||||
NO_SEQUENCE = 0b0000
|
|
||||||
NEG_SEQUENCE = 0b0010
|
NO_SERIALIZATION = 0b0000
|
||||||
JSON_SERIALIZATION = 0b0001
|
JSON = 0b0001
|
||||||
GZIP_COMPRESSION = 0b0001
|
THRIFT = 0b0011
|
||||||
PROTOCOL_VERSION = 0b0001
|
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):
|
class ASRProvider(ASRProviderBase):
|
||||||
def __init__(self, config, delete_audio_file):
|
def __init__(self, config: dict, delete_audio_file: bool):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.interface_type = InterfaceType.STREAM
|
self.interface_type = InterfaceType.NON_STREAM
|
||||||
self.config = config
|
self.appid = config.get("appid")
|
||||||
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.cluster = config.get("cluster")
|
||||||
self.access_token = config.get("access_token")
|
self.access_token = config.get("access_token")
|
||||||
self.boosting_table_name = config.get("boosting_table_name", "")
|
self.boosting_table_name = config.get("boosting_table_name", "")
|
||||||
self.correct_table_name = config.get("correct_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.delete_audio_file = delete_audio_file
|
||||||
|
|
||||||
self.ws_url = "wss://openspeech.bytedance.com/api/v2/asr"
|
self.host = "openspeech.bytedance.com"
|
||||||
self.uid = config.get("uid", "streaming_asr_service")
|
self.ws_url = f"wss://{self.host}/api/v2/asr"
|
||||||
self.workflow = config.get(
|
self.success_code = 1000
|
||||||
"workflow", "audio_in,resample,partition,vad,fe,decode,itn,nlu_punctuate"
|
self.seg_duration = 15000
|
||||||
)
|
|
||||||
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重写父类的方法--开始
|
os.makedirs(self.output_dir, exist_ok=True)
|
||||||
###################################################################################
|
|
||||||
async def open_audio_channels(self, conn):
|
|
||||||
await super().open_audio_channels(conn)
|
|
||||||
|
|
||||||
async with self._session_lock:
|
@staticmethod
|
||||||
# 如果正在重连,等待重连完成
|
def _generate_header(
|
||||||
if self.is_reconnecting:
|
message_type=CLIENT_FULL_REQUEST, message_type_specific_flags=NO_SEQUENCE
|
||||||
logger.bind(tag=TAG).info("等待当前重连完成...")
|
) -> bytearray:
|
||||||
await self._session_close_event.wait()
|
"""Generate protocol header."""
|
||||||
self._session_close_event.clear()
|
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
|
||||||
|
|
||||||
# 如果已有会话未结束,先关闭它
|
def _construct_request(self, reqid) -> dict:
|
||||||
if self._session_started and not self._session_finished:
|
"""Construct the request payload."""
|
||||||
logger.bind(tag=TAG).warning(
|
return {
|
||||||
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": {
|
"app": {
|
||||||
"appid": self.appid,
|
"appid": f"{self.appid}",
|
||||||
"cluster": self.cluster,
|
"cluster": self.cluster,
|
||||||
"token": self.access_token,
|
"token": self.access_token,
|
||||||
},
|
},
|
||||||
"user": {"uid": self.uid},
|
"user": {
|
||||||
|
"uid": str(uuid.uuid4()),
|
||||||
|
},
|
||||||
"request": {
|
"request": {
|
||||||
"reqid": reqid,
|
"reqid": reqid,
|
||||||
"workflow": self.workflow,
|
"show_utterances": False,
|
||||||
"show_utterances": True,
|
|
||||||
"result_type": self.result_type,
|
|
||||||
"sequence": 1,
|
"sequence": 1,
|
||||||
"boosting_table_name": self.boosting_table_name,
|
"boosting_table_name": self.boosting_table_name,
|
||||||
"correct_table_name": self.correct_table_name,
|
"correct_table_name": self.correct_table_name,
|
||||||
},
|
},
|
||||||
"audio": {
|
"audio": {
|
||||||
"format": self.format,
|
"format": "raw",
|
||||||
"codec": self.codec,
|
"rate": 16000,
|
||||||
"rate": self.rate,
|
"language": "zh-CN",
|
||||||
"language": self.language,
|
"bits": 16,
|
||||||
"bits": self.bits,
|
"channel": 1,
|
||||||
"channel": self.channel,
|
"codec": "raw",
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
return req
|
|
||||||
|
|
||||||
def token_auth(self):
|
async def _send_request(
|
||||||
return {"Authorization": f"Bearer; {self.access_token}"}
|
self, audio_data: List[bytes], segment_size: int
|
||||||
|
) -> Optional[str]:
|
||||||
def generate_header(
|
"""Send request to Volcano ASR service."""
|
||||||
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:
|
try:
|
||||||
with self.thread_lock:
|
auth_header = {"Authorization": "Bearer; {}".format(self.access_token)}
|
||||||
if self.conn is None or self.conn.loop is None:
|
async with websockets.connect(
|
||||||
logger.bind(tag=TAG).error(
|
self.ws_url, additional_headers=auth_header
|
||||||
"无法启动ASR监听线程:conn或loop未初始化"
|
) as websocket:
|
||||||
)
|
# Prepare request data
|
||||||
return
|
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:
|
# Send header and metadata
|
||||||
logger.bind(tag=TAG).info("开始启动ASR监听...")
|
# full_client_request
|
||||||
asyncio.run_coroutine_threadsafe(
|
await websocket.send(full_client_request)
|
||||||
self._forward_asr_results(), loop=self.conn.loop
|
res = await websocket.recv()
|
||||||
)
|
result = parse_response(res)
|
||||||
logger.bind(tag=TAG).info("ASR监听已启动")
|
if (
|
||||||
except Exception as e:
|
"payload_msg" in result
|
||||||
logger.bind(tag=TAG).error(f"启动ASR监听线程失败: {e}")
|
and result["payload_msg"]["code"] != self.success_code
|
||||||
except Exception as e:
|
):
|
||||||
logger.bind(tag=TAG).error(f"ASR监听线程发生未预期的错误: {e}")
|
logger.bind(tag=TAG).error(f"ASR error: {result}")
|
||||||
|
return None
|
||||||
|
|
||||||
async def _forward_asr_results(self):
|
for seq, (chunk, last) in enumerate(
|
||||||
try:
|
self.slice_data(audio_data, segment_size), 1
|
||||||
while not self.conn.stop_event.is_set():
|
):
|
||||||
try:
|
if last:
|
||||||
if self.asr_ws is None:
|
audio_only_request = self._generate_header(
|
||||||
# 检查是否需要重连
|
message_type=CLIENT_AUDIO_ONLY_REQUEST,
|
||||||
async with self.reconnect_lock:
|
message_type_specific_flags=NEG_SEQUENCE,
|
||||||
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:
|
else:
|
||||||
try:
|
audio_only_request = self._generate_header(
|
||||||
await self.asr_ws.close()
|
message_type=CLIENT_AUDIO_ONLY_REQUEST
|
||||||
except Exception as e:
|
)
|
||||||
logger.bind(tag=TAG).warning(
|
payload_bytes = gzip.compress(chunk)
|
||||||
f"关闭旧连接时发生错误: {e}"
|
audio_only_request.extend(
|
||||||
)
|
(len(payload_bytes)).to_bytes(4, "big")
|
||||||
finally:
|
) # payload size(4 bytes)
|
||||||
self.asr_ws = None
|
audio_only_request.extend(payload_bytes) # payload
|
||||||
continue
|
# Send audio data
|
||||||
|
await websocket.send(audio_only_request)
|
||||||
|
|
||||||
if "payload_msg" in result:
|
# Receive response
|
||||||
if "result" in result["payload_msg"]:
|
response = await websocket.recv()
|
||||||
# 检查是否有utterances并且definite为True
|
result = parse_response(response)
|
||||||
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:
|
if (
|
||||||
logger.bind(tag=TAG).debug("ASR服务连接已关闭,准备重连...")
|
"payload_msg" in result
|
||||||
# 确保关闭旧连接
|
and result["payload_msg"]["code"] == self.success_code
|
||||||
if self.asr_ws is not None:
|
):
|
||||||
try:
|
if len(result["payload_msg"]["result"]) > 0:
|
||||||
await self.asr_ws.close()
|
return result["payload_msg"]["result"][0]["text"]
|
||||||
except Exception as e:
|
return None
|
||||||
logger.bind(tag=TAG).warning(f"关闭旧连接时发生错误: {e}")
|
else:
|
||||||
finally:
|
logger.bind(tag=TAG).error(f"ASR error: {result}")
|
||||||
self.asr_ws = None
|
return 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:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"ASR监听线程发生错误: {e}")
|
logger.bind(tag=TAG).error(f"ASR request failed: {e}", exc_info=True)
|
||||||
# 确保在发生严重错误时也能继续尝试重连
|
return None
|
||||||
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):
|
@staticmethod
|
||||||
result = self.text
|
def slice_data(data: bytes, chunk_size: int) -> (list, bool):
|
||||||
self.text = "" # 清空text
|
|
||||||
return result, None
|
|
||||||
|
|
||||||
def parse_response(self, res: bytes) -> dict:
|
|
||||||
"""
|
"""
|
||||||
解析 ASR 服务返回的二进制响应。
|
slice data
|
||||||
根据协议格式解析头部和 payload,若采用 GZIP 压缩则先解压,再根据 JSON 反序列化。
|
:param data: wav data
|
||||||
|
:param chunk_size: the segment size in one request
|
||||||
|
:return: segment data, last flag
|
||||||
"""
|
"""
|
||||||
protocol_version = res[0] >> 4
|
data_len = len(data)
|
||||||
header_size = res[0] & 0x0F
|
offset = 0
|
||||||
message_type = res[1] >> 4
|
while offset + chunk_size < data_len:
|
||||||
serialization_method = res[2] >> 4
|
yield data[offset : offset + chunk_size], False
|
||||||
message_compression = res[2] & 0x0F
|
offset += chunk_size
|
||||||
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:
|
else:
|
||||||
payload_msg = payload_msg.decode("utf-8")
|
yield data[offset:data_len], True
|
||||||
result["payload_msg"] = payload_msg
|
|
||||||
result["payload_size"] = payload_size
|
|
||||||
|
|
||||||
# 错误码处理
|
async def speech_to_text(
|
||||||
if "code" in result:
|
self, opus_data: List[bytes], session_id: str
|
||||||
error_code = result["code"]
|
) -> Tuple[Optional[str], Optional[str]]:
|
||||||
error_message = ""
|
"""将语音数据转换为文本"""
|
||||||
|
|
||||||
if error_code == 1000:
|
file_path = None
|
||||||
error_message = "成功"
|
try:
|
||||||
elif error_code == 1001:
|
# 合并所有opus数据包
|
||||||
error_message = "请求参数无效:请求参数缺失必需字段/字段值无效/重复请求"
|
if self.audio_format == "pcm":
|
||||||
elif error_code == 1002:
|
pcm_data = opus_data
|
||||||
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:
|
else:
|
||||||
error_message = "未知错误:未归类错误"
|
pcm_data = self.decode_opus(opus_data)
|
||||||
|
combined_pcm_data = b"".join(pcm_data)
|
||||||
|
|
||||||
logger.bind(tag=TAG).debug(
|
# 判断是否保存为WAV文件
|
||||||
f"ASR错误: {error_message} (错误码: {error_code})"
|
if self.delete_audio_file:
|
||||||
)
|
pass
|
||||||
|
else:
|
||||||
|
file_path = self.save_audio_to_file(pcm_data, session_id)
|
||||||
|
|
||||||
# 如果是识别相关错误,标记需要重连
|
# 直接使用PCM数据
|
||||||
if error_code >= 1020 or error_code == 1001:
|
# 计算分段大小 (单声道, 16bit, 16kHz采样率)
|
||||||
result["need_reconnect"] = True
|
size_per_sec = 1 * 2 * 16000 # nchannels * sampwidth * framerate
|
||||||
|
segment_size = int(size_per_sec * self.seg_duration / 1000)
|
||||||
|
|
||||||
return result
|
# 语音识别
|
||||||
|
start_time = time.time()
|
||||||
async def close_session(self):
|
text = await self._send_request(combined_pcm_data, segment_size)
|
||||||
"""关闭当前会话"""
|
if text:
|
||||||
async with self._session_lock:
|
logger.bind(tag=TAG).debug(
|
||||||
if not self._session_started:
|
f"语音识别耗时: {time.time() - start_time:.3f}s | 结果: {text}"
|
||||||
logger.bind(tag=TAG).warning("尝试关闭未开始的会话")
|
|
||||||
return
|
|
||||||
|
|
||||||
if self._session_finished:
|
|
||||||
logger.bind(tag=TAG).warning(
|
|
||||||
f"会话 {self._current_session_id} 已经关闭"
|
|
||||||
)
|
)
|
||||||
return
|
return text, file_path
|
||||||
|
return "", file_path
|
||||||
|
|
||||||
try:
|
except Exception as e:
|
||||||
if self.asr_ws is not None:
|
logger.bind(tag=TAG).error(f"语音识别失败: {e}", exc_info=True)
|
||||||
await self.asr_ws.close()
|
return "", file_path
|
||||||
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()
|
|
||||||
|
|||||||
@@ -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()
|
||||||
@@ -52,6 +52,7 @@ class TTSProviderBase(ABC):
|
|||||||
)
|
)
|
||||||
self.first_sentence_punctuations = (
|
self.first_sentence_punctuations = (
|
||||||
",",
|
",",
|
||||||
|
"~",
|
||||||
"~",
|
"~",
|
||||||
"、",
|
"、",
|
||||||
",",
|
",",
|
||||||
|
|||||||
Reference in New Issue
Block a user