Files
aistream-test/python/frontend_ws.py
T
2025-12-03 00:27:57 +08:00

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)}")