diff --git a/main/xiaozhi-server/core/providers/asr/base.py b/main/xiaozhi-server/core/providers/asr/base.py index 671484e9..6c6969fd 100644 --- a/main/xiaozhi-server/core/providers/asr/base.py +++ b/main/xiaozhi-server/core/providers/asr/base.py @@ -9,7 +9,6 @@ import asyncio import traceback import threading import opuslib_next -import concurrent.futures from abc import ABC, abstractmethod from config.logger import setup_logging from typing import Optional, Tuple, List @@ -79,98 +78,63 @@ class ASRProviderBase(ABC): """并行处理ASR和声纹识别""" try: total_start_time = time.monotonic() - + # 准备音频数据 if conn.audio_format == "pcm": pcm_data = asr_audio_task else: pcm_data = self.decode_opus(asr_audio_task) - + combined_pcm_data = b"".join(pcm_data) - + # 预先准备WAV数据 wav_data = None if conn.voiceprint_provider and combined_pcm_data: wav_data = self._pcm_to_wav(combined_pcm_data) - + # 定义ASR任务 - def run_asr(): - start_time = time.monotonic() - try: - loop = asyncio.new_event_loop() - asyncio.set_event_loop(loop) - try: - result = loop.run_until_complete( - self.speech_to_text(asr_audio_task, conn.session_id, conn.audio_format) - ) - end_time = time.monotonic() - logger.bind(tag=TAG).debug(f"ASR耗时: {end_time - start_time:.3f}s") - return result - finally: - loop.close() - except Exception as e: - end_time = time.monotonic() - logger.bind(tag=TAG).error(f"ASR失败: {e}") - return ("", None) - - # 定义声纹识别任务 - def run_voiceprint(): - if not wav_data: - return None - try: - loop = asyncio.new_event_loop() - asyncio.set_event_loop(loop) - try: - # 使用连接的声纹识别提供者 - result = loop.run_until_complete( - conn.voiceprint_provider.identify_speaker(wav_data, conn.session_id) - ) - return result - finally: - loop.close() - except Exception as e: - logger.bind(tag=TAG).error(f"声纹识别失败: {e}") - return None - - # 使用线程池执行器并行运行 - with concurrent.futures.ThreadPoolExecutor(max_workers=2) as thread_executor: - asr_future = thread_executor.submit(run_asr) - - if conn.voiceprint_provider and wav_data: - voiceprint_future = thread_executor.submit(run_voiceprint) - - # 等待两个线程都完成 - asr_result = asr_future.result(timeout=15) - voiceprint_result = voiceprint_future.result(timeout=15) - - results = {"asr": asr_result, "voiceprint": voiceprint_result} - else: - asr_result = asr_future.result(timeout=15) - results = {"asr": asr_result, "voiceprint": None} - - - # 处理结果 - raw_text, _ = results.get("asr", ("", None)) - speaker_name = results.get("voiceprint", None) - - # 记录识别结果 + asr_task = self.speech_to_text(asr_audio_task, conn.session_id, conn.audio_format) + + if conn.voiceprint_provider and wav_data: + voiceprint_task = conn.voiceprint_provider.identify_speaker(wav_data, conn.session_id) + # 并发等待两个结果 + asr_result, voiceprint_result = await asyncio.gather( + asr_task, voiceprint_task, return_exceptions=True + ) + else: + asr_result = await asr_task + voiceprint_result = None + + # 记录识别结果 - 检查是否为异常 + if isinstance(asr_result, Exception): + logger.bind(tag=TAG).error(f"ASR识别失败: {asr_result}") + raw_text = "" + else: + raw_text, _ = asr_result + + if isinstance(voiceprint_result, Exception): + logger.bind(tag=TAG).error(f"声纹识别失败: {voiceprint_result}") + speaker_name = "" + else: + speaker_name = voiceprint_result + if raw_text: logger.bind(tag=TAG).info(f"识别文本: {raw_text}") if speaker_name: logger.bind(tag=TAG).info(f"识别说话人: {speaker_name}") - + # 性能监控 total_time = time.monotonic() - total_start_time logger.bind(tag=TAG).debug(f"总处理耗时: {total_time:.3f}s") - + # 检查文本长度 text_len, _ = remove_punctuation_and_length(raw_text) self.stop_ws_connection() - + if text_len > 0: # 构建包含说话人信息的JSON字符串 enhanced_text = self._build_enhanced_text(raw_text, speaker_name) - + # 使用自定义模块进行上报 await startToChat(conn, enhanced_text) enqueue_asr_report(conn, enhanced_text, asr_audio_task)