feat: 将百炼websocket的连接逻辑封装到文件 BaiLianVoiceRecogniseSocket中

This commit is contained in:
martsforever
2025-11-02 19:29:40 +08:00
parent d3d6e1afc6
commit fae4f7f24e
2 changed files with 79 additions and 96 deletions
@@ -10,6 +10,7 @@ from fastapi import FastAPI, WebSocket, Query
from starlette.websockets import WebSocketState from starlette.websockets import WebSocketState
from app.config.env import env from app.config.env import env
from app.utils.bailian_voice_recognise.BaiLianVoiceRecogniseSocket import BaiLianVoiceRecogniseSocket
def add_websocket_voice_recognise_route(app: FastAPI): def add_websocket_voice_recognise_route(app: FastAPI):
@@ -19,117 +20,80 @@ def add_websocket_voice_recognise_route(app: FastAPI):
# 接收连接 # 接收连接
await receive_socket.accept() await receive_socket.accept()
# 百炼WebSocket配置 async def send_text(data: dict):
API_KEY = env.llm_key_bailian await receive_socket.send_text(json.dumps(data, ensure_ascii=False))
QWEN_MODEL = "qwen3-asr-flash-realtime"
baseUrl = "wss://dashscope.aliyuncs.com/api-ws/v1/realtime"
url = f"{baseUrl}?model={QWEN_MODEL}"
headers = { async def on_speech_started(item_id: str):
"Authorization": f"Bearer {API_KEY}", await send_text({"type": "speech_started", "item_id": item_id})
"OpenAI-Beta": "realtime=v1"
}
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: 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(): async def waiting_msg_from_front():
"""将前端PCM数据转发到百炼""" """
将前端PCM数据转发到百炼
"""
try: try:
while True: while True:
# 接收前端PCM数据(二进制格式) int16array_string = await receive_socket.receive_text()
pcm_data_string = await receive_socket.receive_text() int16array = json.loads(int16array_string)
pcm_data_list = json.loads(pcm_data_string)
# print("pcm_data", pcm_data_list)
# print(f"接收到PCM数组数据,长度: {len(pcm_data_list)}")
# 将 int16 列表打包成 bytes # 将 int16 列表打包成 bytes
pcm_bytes = struct.pack('<' + 'h' * len(pcm_data_list), *pcm_data_list) pcm_bytes = struct.pack('<' + 'h' * len(int16array), *int16array)
# 将二进制PCM数据转换为base64编码 # 将二进制PCM数据转换为base64编码
encoded_data = base64.b64encode(pcm_bytes).decode('utf-8') encoded_data = base64.b64encode(pcm_bytes).decode('utf-8')
# 构造发送到百炼的事件 # 构造发送到百炼的事件
audio_event = { audio_event = {
"event_id": f"event_{int(time.time() * 1000)}", "event_id": f"event_{int(time.time() * 1000)}",
"type": "input_audio_buffer.append", "type": "input_audio_buffer.append",
"audio": encoded_data "audio": encoded_data
} }
await bailian_socket.send(json.dumps(audio_event)) await bvrs.socket.send(json.dumps(audio_event))
# print("已转发音频数据到百炼", f"PCM二进制数据,长度: {len(pcm_bytes)} 字节")
except Exception as e: except Exception as e:
if "CloseCode.NO_STATUS_RCVD" in str(e): if "CloseCode.NO_STATUS_RCVD" in str(e):
print("前端连接已断开(无状态码)") print("前端连接已断开(无状态码)")
else: else:
print(f"等待前端的消息出错: {e}") print(f"等待前端的消息出错: {e}")
finally:
async def waiting_msg_from_bailian(): await clear_socket()
"""将百炼识别结果转发到客户端"""
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}")
# 并发运行两个任务 # 并发运行两个任务
await asyncio.gather( await asyncio.gather(
waiting_msg_from_front(), waiting_msg_from_front(),
waiting_msg_from_bailian(), bvrs.waiting_message_from_socket(),
return_exceptions=True return_exceptions=True
) )
except Exception as e: except Exception as e:
print(f"WebSocket连接错误: {e}") print(f"WebSocket连接错误: {e}")
finally: raise e
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)
@@ -1,5 +1,6 @@
import json import json
from typing import Callable, TypedDict from traceback import print_stack
from typing import Callable, TypedDict, Awaitable
import websockets import websockets
from websockets.client import ClientConnection from websockets.client import ClientConnection
@@ -33,13 +34,15 @@ class BaiLianVoiceRecogniseSocket:
enable_server_vad: bool = True, enable_server_vad: bool = True,
# 监听开始说话动作,参数为item_id # 监听开始说话动作,参数为item_id
on_speech_started: Callable[[str], None] | None = None, on_speech_started: Callable[[str], Awaitable[None]] | None = None,
# 监听正在说话的内容,参数为item_id以及本次说话叠加的完整内容 # 监听正在说话的内容,参数为item_id以及本次说话叠加的完整内容
on_speech_content: Callable[[str, str], None] | None = None, on_speech_content: Callable[[str, str], Awaitable[None]] | None = None,
# 监听说话结束动作,参数为item_id # 监听说话结束动作,参数为item_id
on_speech_stopped: Callable[[str], None] | None = None, on_speech_stopped: Callable[[str], Awaitable[None]] | None = None,
# 监听说话完毕动作,参数为item_id以及本次说话万恒内容 # 监听说话完毕动作,参数为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接口地址,以及可选的附加头信息 构造函数,会自动设置连接百炼实时语音识别websocket接口所需要的秘钥,模型名称,websocket接口地址,以及可选的附加头信息
@@ -54,6 +57,7 @@ class BaiLianVoiceRecogniseSocket:
self.on_speech_content = on_speech_content self.on_speech_content = on_speech_content
self.on_speech_stopped = on_speech_stopped self.on_speech_stopped = on_speech_stopped
self.on_speech_completed = on_speech_completed self.on_speech_completed = on_speech_completed
self.on_connect_error = on_connect_error
socket: ClientConnection | None = None socket: ClientConnection | None = None
self.socket = socket self.socket = socket
@@ -119,32 +123,47 @@ class BaiLianVoiceRecogniseSocket:
} }
} }
if self.enable_server_vad: if self.enable_server_vad:
self.log(f"send session.update event: {json.dumps(event_vad, indent=2)}") await self.socket.send(json.dumps(event_vad))
self.socket.send(json.dumps(event_vad))
else: else:
self.log(f"send session.update event: {json.dumps(event_manual, indent=2)}") await self.socket.send(json.dumps(event_manual))
self.socket.send(json.dumps(event_manual))
async def waiting_message_from_socket(self): async def waiting_message_from_socket(self):
"""将百炼识别结果转发到客户端""" """将百炼识别结果转发到客户端"""
try: try:
async for message in self.socket: async for message in self.socket:
data = json.loads(message) data = json.loads(message)
print("data", data)
message_data_type = data.get('type', None) message_data_type = data.get('type', None)
message_data_item_id = data.get('item_id', None) message_data_item_id = data.get('item_id', None)
if message_data_type == "input_audio_buffer.speech_started": 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": 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": 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": 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: except Exception as e:
print_stack(e)
print(f"等待百炼的消息出错: {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): class BaiLianRecogniseMessage(TypedDict):