diff --git a/app/controller/add_websocket_voice_recognise_route.py b/app/controller/add_websocket_voice_recognise_route.py index c7c650c..69923fe 100644 --- a/app/controller/add_websocket_voice_recognise_route.py +++ b/app/controller/add_websocket_voice_recognise_route.py @@ -10,6 +10,7 @@ from fastapi import FastAPI, WebSocket, Query from starlette.websockets import WebSocketState from app.config.env import env +from app.utils.bailian_voice_recognise.BaiLianVoiceRecogniseSocket import BaiLianVoiceRecogniseSocket def add_websocket_voice_recognise_route(app: FastAPI): @@ -19,117 +20,80 @@ def add_websocket_voice_recognise_route(app: FastAPI): # 接收连接 await receive_socket.accept() - # 百炼WebSocket配置 - API_KEY = env.llm_key_bailian - QWEN_MODEL = "qwen3-asr-flash-realtime" - baseUrl = "wss://dashscope.aliyuncs.com/api-ws/v1/realtime" - url = f"{baseUrl}?model={QWEN_MODEL}" + async def send_text(data: dict): + await receive_socket.send_text(json.dumps(data, ensure_ascii=False)) - headers = { - "Authorization": f"Bearer {API_KEY}", - "OpenAI-Beta": "realtime=v1" - } + async def on_speech_started(item_id: str): + await send_text({"type": "speech_started", "item_id": item_id}) - bailian_socket = None + async def on_speech_content(item_id: str, content: str): + await send_text({"type": "speech_content", "item_id": item_id, "content": content}) + + async def on_speech_stopped(item_id: str): + await send_text({"type": "speech_stopped", "item_id": item_id}) + + async def on_speech_completed(item_id: str, content: str): + await send_text({"type": "speech_completed", "item_id": item_id, "content": content}) + + async def on_connect_error(error: Exception): + await clear_socket() + + bvrs = BaiLianVoiceRecogniseSocket( + on_speech_started=on_speech_started, + on_speech_content=on_speech_content, + on_speech_stopped=on_speech_stopped, + on_speech_completed=on_speech_completed, + on_connect_error=on_connect_error, + ) + + await bvrs.connect() + + async def clear_socket(): + """ + 清理掉百炼以及前端的websocket连接 + """ + await bvrs.close() + if receive_socket.application_state != WebSocketState.DISCONNECTED: + print("关闭前端websocket") + await receive_socket.close() + print("前端websocket已经关闭", user_id) try: - # 连接百炼WebSocket - bailian_socket = await websockets.connect(url, additional_headers=headers) - print("已连接到百炼WebSocket") - - # 发送会话更新事件 - session_event = { - "event_id": "event_123", - "type": "session.update", - "session": { - "modalities": ["text"], - "input_audio_format": "pcm", - "sample_rate": 16000, - "input_audio_transcription": { - "language": "zh" - }, - "turn_detection": { - "type": "server_vad", - "threshold": 0.2, - "silence_duration_ms": 800 - } - } - } - await bailian_socket.send(json.dumps(session_event)) - - # 创建双向转发任务 async def waiting_msg_from_front(): - """将前端PCM数据转发到百炼""" + """ + 将前端PCM数据转发到百炼 + """ try: while True: - # 接收前端PCM数据(二进制格式) - pcm_data_string = await receive_socket.receive_text() - pcm_data_list = json.loads(pcm_data_string) - # print("pcm_data", pcm_data_list) - # print(f"接收到PCM数组数据,长度: {len(pcm_data_list)}") - + int16array_string = await receive_socket.receive_text() + int16array = json.loads(int16array_string) # 将 int16 列表打包成 bytes - pcm_bytes = struct.pack('<' + 'h' * len(pcm_data_list), *pcm_data_list) + pcm_bytes = struct.pack('<' + 'h' * len(int16array), *int16array) # 将二进制PCM数据转换为base64编码 encoded_data = base64.b64encode(pcm_bytes).decode('utf-8') - # 构造发送到百炼的事件 audio_event = { "event_id": f"event_{int(time.time() * 1000)}", "type": "input_audio_buffer.append", "audio": encoded_data } - await bailian_socket.send(json.dumps(audio_event)) - # print("已转发音频数据到百炼", f"PCM二进制数据,长度: {len(pcm_bytes)} 字节") + await bvrs.socket.send(json.dumps(audio_event)) except Exception as e: if "CloseCode.NO_STATUS_RCVD" in str(e): print("前端连接已断开(无状态码)") else: print(f"等待前端的消息出错: {e}") - - async def waiting_msg_from_bailian(): - """将百炼识别结果转发到客户端""" - try: - async for message in bailian_socket: - data = json.loads(message) - print("收到百炼消息", data.get("type")) - - # 将识别结果发送给前端 - await receive_socket.send_text(json.dumps(data, ensure_ascii=False)) - - # 如果收到最终转录结果,可以单独处理 - if data.get("type") == "conversation.item.input_audio_transcription.completed": - transcript = data.get("transcript", "") - print(f"最终转录结果: {transcript}") - # 可以发送专门的转录消息给前端 - await receive_socket.send_text(json.dumps({ - "type": "final_transcript", - "text": transcript - }, ensure_ascii=False)) - - except Exception as e: - print(f"等待百炼的消息出错: {e}") + finally: + await clear_socket() # 并发运行两个任务 await asyncio.gather( waiting_msg_from_front(), - waiting_msg_from_bailian(), + bvrs.waiting_message_from_socket(), return_exceptions=True ) except Exception as e: print(f"WebSocket连接错误: {e}") - finally: - print("finally ===>>>>>>>>>>>>>>>>>>>>>>>>>>>") - # 清理连接 - if bailian_socket: - try: - await bailian_socket.close() - except: - pass # 忽略关闭时的任何错误 - - if receive_socket.application_state != WebSocketState.DISCONNECTED: - print("关闭websocket") - await receive_socket.close() - print("websocket已经关闭", user_id) + raise e diff --git a/app/utils/bailian_voice_recognise/BaiLianVoiceRecogniseSocket.py b/app/utils/bailian_voice_recognise/BaiLianVoiceRecogniseSocket.py index 962e251..12c81c5 100644 --- a/app/utils/bailian_voice_recognise/BaiLianVoiceRecogniseSocket.py +++ b/app/utils/bailian_voice_recognise/BaiLianVoiceRecogniseSocket.py @@ -1,5 +1,6 @@ import json -from typing import Callable, TypedDict +from traceback import print_stack +from typing import Callable, TypedDict, Awaitable import websockets from websockets.client import ClientConnection @@ -33,13 +34,15 @@ class BaiLianVoiceRecogniseSocket: enable_server_vad: bool = True, # 监听开始说话动作,参数为item_id - on_speech_started: Callable[[str], None] | None = None, + on_speech_started: Callable[[str], Awaitable[None]] | None = None, # 监听正在说话的内容,参数为item_id以及本次说话叠加的完整内容 - on_speech_content: Callable[[str, str], None] | None = None, + on_speech_content: Callable[[str, str], Awaitable[None]] | None = None, # 监听说话结束动作,参数为item_id - on_speech_stopped: Callable[[str], None] | None = None, + on_speech_stopped: Callable[[str], Awaitable[None]] | None = None, # 监听说话完毕动作,参数为item_id以及本次说话万恒内容 - on_speech_completed: Callable[[str, str], None] | None = None + on_speech_completed: Callable[[str, str], Awaitable[None]] | None = None, + # 连接错误处理 + on_connect_error: Callable[[Exception], Awaitable[None]] | None = None, ): """ 构造函数,会自动设置连接百炼实时语音识别websocket接口所需要的秘钥,模型名称,websocket接口地址,以及可选的附加头信息 @@ -54,6 +57,7 @@ class BaiLianVoiceRecogniseSocket: self.on_speech_content = on_speech_content self.on_speech_stopped = on_speech_stopped self.on_speech_completed = on_speech_completed + self.on_connect_error = on_connect_error socket: ClientConnection | None = None self.socket = socket @@ -119,32 +123,47 @@ class BaiLianVoiceRecogniseSocket: } } if self.enable_server_vad: - self.log(f"send session.update event: {json.dumps(event_vad, indent=2)}") - self.socket.send(json.dumps(event_vad)) + await self.socket.send(json.dumps(event_vad)) else: - self.log(f"send session.update event: {json.dumps(event_manual, indent=2)}") - self.socket.send(json.dumps(event_manual)) + await self.socket.send(json.dumps(event_manual)) async def waiting_message_from_socket(self): """将百炼识别结果转发到客户端""" try: async for message in self.socket: data = json.loads(message) + print("data", data) message_data_type = data.get('type', None) message_data_item_id = data.get('item_id', None) if message_data_type == "input_audio_buffer.speech_started": - self.on_speech_started(message_data_item_id) + await self.on_speech_started(message_data_item_id) if message_data_type == "conversation.item.input_audio_transcription.text": - self.on_speech_content(message_data_item_id, data.get('stash', "")) + await self.on_speech_content(message_data_item_id, data.get('stash', "")) if message_data_type == "input_audio_buffer.speech_stopped": - self.on_speech_stopped(message_data_item_id) + await self.on_speech_stopped(message_data_item_id) if message_data_type == "conversation.item.input_audio_transcription.completed": - self.on_speech_completed(message_data_item_id, data.get('transcript', "")) + await self.on_speech_completed(message_data_item_id, data.get('transcript', "")) except Exception as e: + print_stack(e) print(f"等待百炼的消息出错: {e}") + await self.on_connect_error(e) + + async def close(self): + """ + 关闭websocket连接 + """ + if self.socket: + try: + print("准备关闭百炼websocket") + await self.socket.close() + print("百炼websocket已经关闭") + except: + pass + finally: + self.socket = None class BaiLianRecogniseMessage(TypedDict):