293 lines
11 KiB
Python
293 lines
11 KiB
Python
# frontend_ws.py
|
|
from fastapi import WebSocket, WebSocketDisconnect
|
|
from session_manager import create_session, close_session
|
|
|
|
# 导入大模型模块
|
|
from llm_client import call_llm, LLMConversation, llm_client
|
|
from typing import Optional, Dict, Callable, Awaitable, List
|
|
from dataclasses import dataclass, field
|
|
import sys
|
|
import json
|
|
import asyncio
|
|
import asyncio
|
|
import json
|
|
import websockets
|
|
from typing import Optional, Callable, Dict, Any, Coroutine
|
|
from dataclasses import dataclass, field
|
|
import uuid
|
|
|
|
import logging
|
|
logger = logging.getLogger(__name__)
|
|
from ws_message_manager import (
|
|
ws_queue_manager,
|
|
ClientMsgType,
|
|
ServerMsgType,
|
|
)
|
|
|
|
from asr_client import (
|
|
init_asr_pool,
|
|
get_idle_asr_connection,
|
|
handle_asr_communication,
|
|
push_audio_data, # 导入音频插入接口
|
|
close_asr_pool
|
|
)
|
|
from tts_client import TTSManager
|
|
active_connections = []
|
|
user_llm_conversations: dict[str, LLMConversation] = {}
|
|
consume_wakeup = asyncio.Event()
|
|
|
|
|
|
async def frontend_websocket_handler(websocket: WebSocket):
|
|
"""前端WebSocket入口处理器(修复 TTS 数据接收问题)"""
|
|
# 1. 握手+创建会话
|
|
await websocket.accept()
|
|
active_connections.append(websocket)
|
|
conn_id = id(websocket)
|
|
user_id = f"user_{conn_id}"
|
|
session = await create_session(websocket)
|
|
logger.info(f"连接 {conn_id} 建立成功")
|
|
|
|
tts_session_id = user_id
|
|
result_queue = asyncio.Queue(maxsize=10000)
|
|
asr_task = None
|
|
asr_conn = None
|
|
|
|
# 初始化用户LLM会话
|
|
if user_id not in user_llm_conversations:
|
|
user_llm_conversations[user_id] = LLMConversation(
|
|
user_id=user_id,
|
|
scene_description="语音识别对话场景"
|
|
)
|
|
llm_conversation = user_llm_conversations[user_id]
|
|
|
|
# ====================== 修复 TTS 核心逻辑 ======================
|
|
# 1. 创建 TTS 客户端
|
|
# 1. 创建TTS管理器
|
|
# tts_manager = TTSManager(ws_url="ws://10.10.10.202:50000/ws/tts")
|
|
tts_manager = TTSManager()
|
|
# 2. 设置结果回调函数(接收完整结果)
|
|
def handle_tts_result(req_id: str, result: Dict[str, Any]):
|
|
"""处理TTS结果回调"""
|
|
# print('处理TTS结果回调 ',result)
|
|
status = result.get("status")
|
|
|
|
if status == "completed":
|
|
audio_data = result.get("audio_data")
|
|
sample_rate = result.get("sample_rate")
|
|
result_queue.put_nowait(audio_data)
|
|
#
|
|
# if audio_data is not None and len(audio_data) > 0:
|
|
# # 转换为PCM数据
|
|
# pcm_data = (audio_data.astype(np.float32) * 32767).astype(np.int16)
|
|
# pcm_bytes = pcm_data.tobytes()
|
|
# # result_queue.put_nowait({"type": 3, "data": pcm_bytes})
|
|
# result_queue.put_nowait(pcm_bytes)
|
|
# print(f"插入时候队列当前大小xxx: {result_queue.qsize()}") # 排查队列是否有数据
|
|
# 保存为PCM文件
|
|
|
|
|
|
# elif status == "error":
|
|
# error_msg = result.get("message")
|
|
# print(f"❌ TTS处理失败 [{req_id[:8]}]: {error_msg}")
|
|
|
|
tts_manager.set_result_callback(handle_tts_result)
|
|
|
|
# 3. 设置是否播放(可选,默认True)
|
|
tts_manager.set_playback_enabled(False) # 设置为False则不播放
|
|
|
|
# 4. 初始化连接
|
|
await tts_manager.initialize()
|
|
|
|
# texts = [
|
|
# "你好,这是第一个排队的TTS请求。你好,这是第一个排队的TTS请求。你好,这是第一个排队的TTS请好",
|
|
# "我是第二个请求,会等第一个处理完再执行。",
|
|
# "第三个请求,支持流式播放和队列管理。",
|
|
# "第四个请求,测试队列的自动消费功能。",
|
|
# "最后一个请求,处理完成后会自动结束。"
|
|
# ]
|
|
# req_ids = []
|
|
# for i, text in enumerate(texts):
|
|
# req_id = await tts_manager.synthesize(text)
|
|
# req_ids.append(req_id)
|
|
# ====================== ASR 结果回调 ======================
|
|
async def asr_result_callback(result: dict):
|
|
"""ASR 结果回调:转发前端 + 调用大模型"""
|
|
try:
|
|
logger.info(f"ASR 识别结果: {result}")
|
|
final_asr_text = result.get("text", "")
|
|
|
|
# 1. 转发 ASR 结果给前端
|
|
# if not websocket.client_state.disconnected:
|
|
# await websocket.send_json({
|
|
# "type": "asr_result",
|
|
# "data": result
|
|
# })
|
|
# result_queue.put_nowait({"type": 1, "data": final_asr_text})
|
|
print(f"插入时候队列当前大小: {result_queue.qsize()}") # 排查队列是否有数据
|
|
# 2. 立即唤醒消费协程(无延迟)
|
|
consume_wakeup.set()
|
|
# 2. 调用大模型(异步)
|
|
if final_asr_text:
|
|
asyncio.create_task(
|
|
call_llm_and_send(
|
|
query=final_asr_text,
|
|
conversation=llm_conversation
|
|
)
|
|
)
|
|
|
|
except Exception as e:
|
|
logger.error(f"ASR 回调执行失败: {str(e)}")
|
|
|
|
# ====================== 大模型流式回调 ======================
|
|
async def llm_stream_callback(chunk: str, conversation_id: str, is_finished: bool):
|
|
"""大模型流式回调(纯异步,无阻塞)"""
|
|
if not chunk:
|
|
return
|
|
# req_id = await tts_manager.synthesize(
|
|
# chunk,
|
|
# mode="预训练音色",
|
|
# sft_spk="中文女",
|
|
# speed=1.0
|
|
# )
|
|
req_id = await tts_manager.synthesize(chunk)
|
|
# 1. 异步插入队列(替代put_nowait,避免队列满时抛异常)
|
|
# 1. 非阻塞插入(队列满则丢弃,优先保证实时性)
|
|
# result_queue.put_nowait(chunk)
|
|
# result_queue.put_nowait({"type": 2, "data": chunk})
|
|
# print(f"📥 插入队列: {chunk}, 队列大小: {result_queue.qsize()}")
|
|
|
|
|
|
|
|
# 3. 让出调度权,确保消费协程执行
|
|
await asyncio.sleep(0)
|
|
|
|
# ====================== 调用大模型 ======================
|
|
async def call_llm_and_send(query: str, conversation: LLMConversation):
|
|
"""调用大模型,流式结果转发前端 + TTS"""
|
|
if not query:
|
|
return
|
|
logger.info(f"调用大模型 - 用户({user_id}): {query}")
|
|
|
|
try:
|
|
conv_id, full_reply = await llm_client.send_message(
|
|
query=query,
|
|
conversation=conversation,
|
|
stream_callback=llm_stream_callback,
|
|
response_mode="streaming"
|
|
)
|
|
logger.info(f"大模型回复完成 - 会话ID: {conv_id}, 完整回复: {full_reply}")
|
|
except Exception as e:
|
|
logger.error(f"大模型调用失败: {str(e)}")
|
|
await websocket.send_json({
|
|
"type": "llm_error",
|
|
"data": {"error": str(e)}
|
|
})
|
|
|
|
# ====================== ASR 处理 ======================
|
|
try:
|
|
# 获取 ASR 连接
|
|
asr_conn = await get_idle_asr_connection()
|
|
if not asr_conn:
|
|
await websocket.send_json({"error": "ASR 服务暂时不可用", "text": ""})
|
|
return
|
|
|
|
# 启动 ASR 协程
|
|
asr_task = asyncio.create_task(
|
|
handle_asr_communication(asr_conn, asr_result_callback)
|
|
)
|
|
# 接收前端数据
|
|
async def recv_frontend_data():
|
|
"""接收前端音频/控制指令"""
|
|
while not asr_conn.stop_event.is_set():
|
|
try:
|
|
if not result_queue.empty():
|
|
await asyncio.sleep(0) # 立即让权
|
|
continue
|
|
raw_bytes = await websocket.receive_bytes()
|
|
success = await push_audio_data(asr_conn, raw_bytes)
|
|
if not success:
|
|
print("音频数据插入失败(队列满/连接失效)")
|
|
|
|
except WebSocketDisconnect:
|
|
logger.info(f"前端 {conn_id} 主动断开连接")
|
|
asr_conn.stop_event.set()
|
|
break
|
|
except Exception as e:
|
|
logger.error(f"接收前端数据失败: {str(e)}")
|
|
asr_conn.stop_event.set()
|
|
await websocket.send_json({"error": f"接收数据失败: {str(e)}"})
|
|
break
|
|
|
|
# 发送 ASR 结果
|
|
async def send_asr_result():
|
|
"""从结果队列发送 ASR 结果到前端(二进制格式)"""
|
|
# while not asr_conn.stop_event.is_set():
|
|
while True:
|
|
try:
|
|
# print('XXXXXXXXXXXXXXX')
|
|
# print(f"消费者队列当前大小: {result_queue.qsize()}") # 排查队列是否有数据
|
|
result = await asyncio.wait_for(result_queue.get(), timeout=0.05)
|
|
# print('YYYYYYYYYYYYYYYYYYY')
|
|
# # 直接发送二进制数据,不进行JSON格式化
|
|
# await websocket.send_json(result)
|
|
await websocket.send_bytes(result)
|
|
|
|
|
|
|
|
except asyncio.TimeoutError:
|
|
continue
|
|
except Exception as e:
|
|
logger.error(f"发送 ASR 结果失败: {str(e)}")
|
|
asr_conn.stop_event.set()
|
|
break
|
|
|
|
|
|
task_send = asyncio.create_task(send_asr_result())
|
|
task_recv = asyncio.create_task(recv_frontend_data())
|
|
|
|
|
|
|
|
|
|
|
|
# 2. 等待任一任务完成(或stop_event触发),而非等待两者都完成
|
|
try:
|
|
# 等待两个任务,只要有一个完成就返回(比如前端断开/发送出错)
|
|
done, pending = await asyncio.wait(
|
|
[task_recv, task_send],
|
|
return_when=asyncio.FIRST_COMPLETED,
|
|
timeout=None # 无限等待,直到有任务完成
|
|
)
|
|
finally:
|
|
# 确保协程正确退出
|
|
asr_conn.stop_event.set()
|
|
# 等待剩余任务完成
|
|
for task in pending:
|
|
task.cancel()
|
|
await asyncio.gather(task_recv, task_send, return_exceptions=True)
|
|
|
|
except Exception as e:
|
|
logger.error(f"WebSocket 处理异常: {str(e)}")
|
|
if asr_conn:
|
|
asr_conn.stop_event.set()
|
|
await websocket.send_json({"error": str(e)})
|
|
|
|
finally:
|
|
# 清理资源
|
|
if asr_conn:
|
|
asr_conn.stop_event.set()
|
|
|
|
if asr_task and not asr_task.done():
|
|
asr_task.cancel()
|
|
try:
|
|
await asr_task
|
|
except asyncio.CancelledError:
|
|
pass
|
|
|
|
# 修复点8:清理 TTS 资源
|
|
# tts_client.unregister_session_callback(tts_session_id)
|
|
# await tts_client.disconnect()
|
|
|
|
if websocket in active_connections:
|
|
active_connections.remove(websocket)
|
|
|
|
logger.info(f"连接 {conn_id} 已清理,当前连接数: {len(active_connections)}") |