343 lines
12 KiB
Python
343 lines
12 KiB
Python
import asyncio
|
||
import threading
|
||
from dashscope.audio.asr import TranslationRecognizerChat
|
||
from dashscope.audio.asr import TranscriptionResult
|
||
from dashscope.audio.asr import TranslationRecognizerCallback
|
||
from config import settings
|
||
from interfaces.asr import ASR
|
||
from implementations.base_service import BaseService
|
||
from services.interrupt_handler import interrupt_handler
|
||
from utils.logger import session_logger
|
||
|
||
|
||
class GummyCallback(TranslationRecognizerCallback):
|
||
def __init__(self, device_id, session_id, asr_instance=None):
|
||
self.device_id = device_id
|
||
self.session_id = session_id
|
||
self.asr_instance = asr_instance
|
||
self.transcript = ""
|
||
self.transcript_event = asyncio.Event()
|
||
self.received_result = False
|
||
self.error_message = None
|
||
self.is_closed = False
|
||
self.session_end = False
|
||
self.reconnect_event = threading.Event()
|
||
self.is_rate_limit_error = False
|
||
|
||
def on_open(self) -> None:
|
||
session_logger.info(self.device_id, self.session_id, "Gummy ASR连接建立成功")
|
||
|
||
def on_event(
|
||
self,
|
||
request_id,
|
||
transcription_result: TranscriptionResult,
|
||
translation_result,
|
||
usage,
|
||
) -> None:
|
||
if transcription_result:
|
||
if transcription_result.is_sentence_end:
|
||
self.transcript = transcription_result.text
|
||
self.received_result = True
|
||
session_logger.info(
|
||
self.device_id,
|
||
self.session_id,
|
||
f"Gummy ASR完整句子: {self.transcript}",
|
||
)
|
||
self.transcript_event.set()
|
||
else:
|
||
self.transcript = transcription_result.text
|
||
|
||
def on_complete(self) -> None:
|
||
session_logger.info(self.device_id, self.session_id, "Gummy ASR识别完成")
|
||
self.session_end = True
|
||
self.transcript_event.set()
|
||
|
||
def on_error(self, result) -> None:
|
||
self.error_message = str(result)
|
||
error_str = str(result).lower()
|
||
self.is_rate_limit_error = "throttling" in error_str or "rate limit" in error_str
|
||
|
||
if self.is_rate_limit_error:
|
||
session_logger.warning(
|
||
self.device_id,
|
||
self.session_id,
|
||
f"Gummy ASR限流错误,触发重连信号: {self.error_message}"
|
||
)
|
||
self.reconnect_event.set()
|
||
if self.asr_instance:
|
||
self.asr_instance.need_reconnect = True
|
||
else:
|
||
session_logger.error(
|
||
self.device_id, self.session_id, f"Gummy ASR错误: {self.error_message}"
|
||
)
|
||
self.transcript_event.set()
|
||
|
||
def on_close(self) -> None:
|
||
session_logger.info(self.device_id, self.session_id, "Gummy ASR连接已关闭")
|
||
self.is_closed = True
|
||
self.transcript_event.set()
|
||
|
||
|
||
class AliyunASR(ASR, BaseService):
|
||
def __init__(self, selected_role=None):
|
||
BaseService.__init__(self, "asr", selected_role)
|
||
self.api_key = settings.aliyun_api_key
|
||
self.vocabulary_id = settings.aliyun_vocabulary_id
|
||
self.connected = False
|
||
self.transcript = ""
|
||
self.device_id = "unknown"
|
||
self.session_id = "unknown"
|
||
self.sample_rate = 16000
|
||
self.audio_buffer = bytearray()
|
||
self.translator = None
|
||
self.callback = None
|
||
self.buffer_size = 3200
|
||
self.max_end_silence = 2000
|
||
self.max_retries = 3
|
||
self.retry_delay = 1.0
|
||
self.need_reconnect = False
|
||
|
||
async def connect(self):
|
||
try:
|
||
if self.connected and self.translator:
|
||
session_logger.info(
|
||
self.device_id, self.session_id, "Gummy ASR连接已存在,复用当前连接"
|
||
)
|
||
return
|
||
|
||
self.transcript = ""
|
||
|
||
self.callback = GummyCallback(self.device_id, self.session_id, self)
|
||
|
||
self.translator = TranslationRecognizerChat(
|
||
model="gummy-chat-v1",
|
||
format="mp3",
|
||
sample_rate=self.sample_rate,
|
||
callback=self.callback,
|
||
transcription_enabled=True,
|
||
source_language="auto",
|
||
semantic_punctuation_enabled=False,
|
||
max_end_silence=self.max_end_silence,
|
||
api_key=self.api_key,
|
||
vocabulary_id=self.vocabulary_id,
|
||
)
|
||
|
||
self.translator.start()
|
||
self.callback.transcript_event.clear()
|
||
self.connected = True
|
||
self.need_reconnect = False
|
||
|
||
except Exception as e:
|
||
session_logger.error(
|
||
self.device_id, self.session_id, f"建立Gummy ASR连接失败: {str(e)}"
|
||
)
|
||
await self._cleanup_connection()
|
||
|
||
async def _cleanup_connection(self):
|
||
if self.translator:
|
||
try:
|
||
self.translator.stop()
|
||
except:
|
||
pass
|
||
self.translator = None
|
||
self.callback = None
|
||
self.connected = False
|
||
|
||
async def _handle_reconnect(self):
|
||
if not self.need_reconnect:
|
||
return True
|
||
|
||
session_logger.warning(
|
||
self.device_id,
|
||
self.session_id,
|
||
"检测到限流错误,准备重连..."
|
||
)
|
||
|
||
await self._cleanup_connection()
|
||
|
||
current_delay = self.retry_delay
|
||
for attempt in range(self.max_retries):
|
||
session_logger.info(
|
||
self.device_id,
|
||
self.session_id,
|
||
f"重连尝试 {attempt + 1}/{self.max_retries},等待{current_delay}秒"
|
||
)
|
||
await asyncio.sleep(current_delay)
|
||
|
||
try:
|
||
await self.connect()
|
||
if self.connected:
|
||
session_logger.info(
|
||
self.device_id,
|
||
self.session_id,
|
||
"重连成功"
|
||
)
|
||
return True
|
||
except Exception as e:
|
||
session_logger.error(
|
||
self.device_id,
|
||
self.session_id,
|
||
f"重连失败: {str(e)}"
|
||
)
|
||
|
||
current_delay *= 2
|
||
self.retry_delay = current_delay
|
||
|
||
session_logger.error(
|
||
self.device_id,
|
||
self.session_id,
|
||
"重连失败,已达到最大重试次数"
|
||
)
|
||
self.need_reconnect = False
|
||
return False
|
||
|
||
async def send_audio(self, audio_data: bytes):
|
||
if self.need_reconnect:
|
||
reconnected = await self._handle_reconnect()
|
||
if not reconnected:
|
||
session_logger.error(self.device_id, self.session_id, "Gummy ASR重连失败,无法发送音频")
|
||
return
|
||
|
||
if not self.connected or not self.translator:
|
||
session_logger.error(self.device_id, self.session_id, "Gummy ASR连接未建立")
|
||
return
|
||
|
||
try:
|
||
self.audio_buffer.extend(audio_data)
|
||
|
||
while len(self.audio_buffer) >= self.buffer_size:
|
||
if self.need_reconnect:
|
||
session_logger.info(
|
||
self.device_id,
|
||
self.session_id,
|
||
"检测到限流错误,暂停发送音频,准备重连"
|
||
)
|
||
reconnected = await self._handle_reconnect()
|
||
if not reconnected:
|
||
return
|
||
if not self.connected or not self.translator:
|
||
return
|
||
continue
|
||
|
||
chunk = bytes(self.audio_buffer[: self.buffer_size])
|
||
self.audio_buffer = self.audio_buffer[self.buffer_size :]
|
||
|
||
if self.device_id and self.session_id:
|
||
session_key = (self.device_id, self.session_id)
|
||
if interrupt_handler.is_interrupted(session_key):
|
||
session_logger.info(
|
||
self.device_id,
|
||
self.session_id,
|
||
"检测到中断,停止发送音频数据",
|
||
)
|
||
return
|
||
|
||
try:
|
||
self.translator.send_audio_frame(chunk)
|
||
except Exception as e:
|
||
session_logger.error(
|
||
self.device_id,
|
||
self.session_id,
|
||
f"发送音频帧到Gummy ASR出错: {str(e)}",
|
||
)
|
||
|
||
except Exception as e:
|
||
session_logger.error(
|
||
self.device_id, self.session_id, f"发送音频到Gummy ASR出错: {str(e)}"
|
||
)
|
||
|
||
async def send_end(self):
|
||
if not self.connected or not self.translator:
|
||
session_logger.error(
|
||
self.device_id, self.session_id, "Gummy ASR连接未建立,无法发送结束信号"
|
||
)
|
||
return
|
||
|
||
try:
|
||
if len(self.audio_buffer) > 0:
|
||
chunk = bytes(self.audio_buffer)
|
||
self.audio_buffer = bytearray()
|
||
self.translator.send_audio_frame(chunk)
|
||
|
||
session_logger.info(
|
||
self.device_id,
|
||
self.session_id,
|
||
"客户端发送结束包,等待Gummy ASR处理完毕",
|
||
)
|
||
|
||
loop = asyncio.get_running_loop()
|
||
await loop.run_in_executor(None, lambda: self._safe_stop_translator())
|
||
|
||
except Exception as e:
|
||
session_logger.error(
|
||
self.device_id,
|
||
self.session_id,
|
||
f"发送结束信号到Gummy ASR出错: {str(e)}",
|
||
)
|
||
|
||
async def receive_results(self):
|
||
if not self.connected:
|
||
session_logger.error(
|
||
self.device_id, self.session_id, "Gummy ASR连接未建立,无法接收结果"
|
||
)
|
||
return
|
||
|
||
try:
|
||
try:
|
||
await asyncio.wait_for(
|
||
self.callback.transcript_event.wait(), timeout=0.5
|
||
)
|
||
except asyncio.TimeoutError:
|
||
session_logger.warning(
|
||
self.device_id,
|
||
self.session_id,
|
||
"接收Gummy ASR结果超时,使用最后的中间结果",
|
||
)
|
||
|
||
if self.callback.received_result or self.callback.transcript:
|
||
self.transcript = self.callback.transcript
|
||
|
||
except Exception as e:
|
||
session_logger.error(
|
||
self.device_id, self.session_id, f"接收Gummy ASR结果出错: {str(e)}"
|
||
)
|
||
|
||
return self.transcript
|
||
|
||
async def close(self):
|
||
if self.connected:
|
||
try:
|
||
if self.translator is not None:
|
||
try:
|
||
self.translator.stop()
|
||
except Exception as e:
|
||
if "has stopped" not in str(e):
|
||
session_logger.warning(
|
||
self.device_id,
|
||
self.session_id,
|
||
f"关闭Gummy translator出错: {str(e)}"
|
||
)
|
||
self.translator = None
|
||
self.connected = False
|
||
session_logger.info(
|
||
self.device_id, self.session_id, "Gummy ASR连接已关闭"
|
||
)
|
||
except Exception as e:
|
||
session_logger.error(
|
||
self.device_id, self.session_id, f"关闭Gummy ASR连接出错: {str(e)}"
|
||
)
|
||
|
||
def _safe_stop_translator(self):
|
||
"""安全停止translator,忽略已停止的异常"""
|
||
try:
|
||
if self.translator:
|
||
self.translator.stop()
|
||
except Exception as e:
|
||
if "has stopped" not in str(e):
|
||
session_logger.warning(
|
||
self.device_id,
|
||
self.session_id,
|
||
f"停止translator时出现非标准错误: {str(e)}",
|
||
)
|
||
return None
|