mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-29 13:03:56 +08:00
Merge branch 'test-mcp' into py_test_mcp
This commit is contained in:
@@ -47,7 +47,7 @@ class ASRProviderBase(ABC):
|
||||
if self.conn.client_voice_stop:
|
||||
asr_audio_task = copy.deepcopy(self.conn.asr_audio)
|
||||
self.conn.asr_audio.clear()
|
||||
self.conn.client_abort = False
|
||||
|
||||
# 音频太短了,无法识别
|
||||
self.conn.reset_vad_states()
|
||||
if len(asr_audio_task) > 15:
|
||||
|
||||
@@ -33,6 +33,21 @@ class ASRProvider(ASRProviderBase):
|
||||
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")
|
||||
@@ -60,7 +75,6 @@ class ASRProvider(ASRProviderBase):
|
||||
self.asr_ws = None
|
||||
self.forward_task = None
|
||||
self.conn = None
|
||||
self.asr_thread = None
|
||||
|
||||
###################################################################################
|
||||
# 豆包流式ASR重写父类的方法--开始
|
||||
@@ -68,10 +82,18 @@ class ASRProvider(ASRProviderBase):
|
||||
async def open_audio_channels(self, conn):
|
||||
await super().open_audio_channels(conn)
|
||||
|
||||
retry_count = 0
|
||||
while retry_count < self.max_retries:
|
||||
try:
|
||||
# 确保关闭旧的连接
|
||||
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()
|
||||
@@ -79,62 +101,94 @@ class ASRProvider(ASRProviderBase):
|
||||
logger.bind(tag=TAG).warning(f"关闭旧连接时发生错误: {e}")
|
||||
finally:
|
||||
self.asr_ws = None
|
||||
self._session_finished = True
|
||||
self._session_close_event.set()
|
||||
|
||||
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,因为服务器可能不支持
|
||||
ping_timeout=None,
|
||||
close_timeout=10,
|
||||
)
|
||||
# 重置会话状态
|
||||
self._current_session_id = str(uuid.uuid4())
|
||||
self._session_started = True
|
||||
self._session_finished = False
|
||||
self.is_reconnecting = True
|
||||
|
||||
# 发送初始化请求
|
||||
request_params = self.construct_request(str(uuid.uuid4()))
|
||||
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)
|
||||
logger.bind(tag=TAG).debug(f"发送初始化请求: {request_params}")
|
||||
except Exception as e:
|
||||
logger.bind(tag=TAG).error(f"发送初始化请求失败: {e}")
|
||||
raise e
|
||||
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,
|
||||
)
|
||||
|
||||
# 等待初始化响应
|
||||
try:
|
||||
init_res = await self.asr_ws.recv()
|
||||
result = self.parse_response(init_res)
|
||||
logger.bind(tag=TAG).info(f"ASR服务初始化响应: {result}")
|
||||
except Exception as e:
|
||||
logger.bind(tag=TAG).error(f"ASR服务初始化失败: {e}")
|
||||
raise e
|
||||
# 发送初始化请求
|
||||
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
|
||||
|
||||
# 启动接收ASR结果的异步任务
|
||||
asr_priority = threading.Thread(
|
||||
target=self._start_monitor_asr_response_thread, daemon=True
|
||||
)
|
||||
asr_priority.start()
|
||||
return
|
||||
# 等待初始化响应
|
||||
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
|
||||
|
||||
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).error(
|
||||
f"WebSocket连接失败,已达到最大重试次数: {e}"
|
||||
)
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.bind(tag=TAG).error(f"WebSocket连接发生未知错误: {e}")
|
||||
raise
|
||||
# 启动接收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):
|
||||
@@ -150,7 +204,7 @@ class ASRProvider(ASRProviderBase):
|
||||
if self.asr_ws:
|
||||
await self.asr_ws.send(audio_request)
|
||||
except Exception as e:
|
||||
logger.bind(tag=TAG).error(f"发送音频数据时发生错误: {e}")
|
||||
logger.bind(tag=TAG).debug(f"发送音频数据时发生错误: {e}")
|
||||
|
||||
###################################################################################
|
||||
# 豆包流式ASR重写父类的方法--结束
|
||||
@@ -246,23 +300,61 @@ class ASRProvider(ASRProviderBase):
|
||||
|
||||
def _start_monitor_asr_response_thread(self):
|
||||
# 初始化链接
|
||||
asyncio.run_coroutine_threadsafe(
|
||||
self._forward_asr_results(), loop=self.conn.loop
|
||||
)
|
||||
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:
|
||||
logger.bind(tag=TAG).info("尝试重新连接ASR服务...")
|
||||
await self.open_audio_channels(self.conn)
|
||||
continue
|
||||
# 检查是否需要重连
|
||||
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(
|
||||
@@ -290,6 +382,7 @@ class ASRProvider(ASRProviderBase):
|
||||
self.text = utterance["text"]
|
||||
await self.handle_voice_stop(None)
|
||||
break
|
||||
|
||||
except websockets.ConnectionClosed:
|
||||
logger.bind(tag=TAG).debug("ASR服务连接已关闭,准备重连...")
|
||||
# 确保关闭旧连接
|
||||
@@ -301,34 +394,16 @@ class ASRProvider(ASRProviderBase):
|
||||
finally:
|
||||
self.asr_ws = None
|
||||
|
||||
retry_count = 0
|
||||
while (
|
||||
retry_count < self.max_retries
|
||||
and not self.conn.stop_event.is_set()
|
||||
):
|
||||
try:
|
||||
logger.bind(tag=TAG).info(
|
||||
f"正在进行第{retry_count + 1}次重连尝试..."
|
||||
)
|
||||
await self.open_audio_channels(self.conn)
|
||||
break
|
||||
except Exception as e:
|
||||
retry_count += 1
|
||||
if retry_count < self.max_retries:
|
||||
logger.bind(tag=TAG).warning(
|
||||
f"重连失败,等待{self.retry_delay}秒后重试: {e}"
|
||||
)
|
||||
await asyncio.sleep(self.retry_delay)
|
||||
else:
|
||||
logger.bind(tag=TAG).error(
|
||||
f"重连失败,已达到最大重试次数: {e}"
|
||||
)
|
||||
await asyncio.sleep(
|
||||
self.retry_delay
|
||||
) # 继续等待,以便后续重试
|
||||
# 等待冷却时间
|
||||
await asyncio.sleep(self.reconnect_cooldown)
|
||||
continue
|
||||
|
||||
except Exception as e:
|
||||
if not self.conn.stop_event.is_set():
|
||||
await asyncio.sleep(2) # 增加重试延迟
|
||||
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}")
|
||||
# 确保在发生严重错误时也能继续尝试重连
|
||||
@@ -422,8 +497,38 @@ class ASRProvider(ASRProviderBase):
|
||||
f"ASR错误: {error_message} (错误码: {error_code})"
|
||||
)
|
||||
|
||||
# 如果是识别相关错误(>=1020),标记需要重连
|
||||
if error_code >= 1020:
|
||||
# 如果是识别相关错误,标记需要重连
|
||||
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()
|
||||
|
||||
Reference in New Issue
Block a user