# 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 import numpy as np from typing import Optional, Dict, Callable, Awaitable, List from dataclasses import dataclass, field import sys import json import asyncio 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 CosyVoiceTTSSocketClient active_connections = [] user_llm_conversations: dict[str, LLMConversation] = {} async def frontend_websocket_handler(websocket: WebSocket): """前端WebSocket入口处理器(保留Session,ASR/TTS无状态)""" # 1. 握手+创建会话(核心:Session 管理连接上下文) await websocket.accept() active_connections.append(websocket) conn_id = id(websocket) user_id = f"user_{id(websocket)}" # 实际场景替换为真实用户ID session = await create_session(websocket) # 创建会话(无用户逻辑) logger.info(f"连接 {conn_id} 建立成功") # await ws_queue_manager.create_queue(websocket) tts_session_id = user_id result_queue = asyncio.Queue(maxsize=100) 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] # 1. 创建客户端 client = CosyVoiceTTSSocketClient( ws_url="ws://10.10.10.202:50000/ws/tts", max_queue_size=500 # 最大队列长度 ) # 2. 注册回调(按需自定义) client.on_audio_chunk = lambda req_id, chunk: print(f"收到[{req_id[:8]}]音频块,长度: {len(chunk)}") # 3. 建立连接(自动启动队列消费) await client.connect() # ---------------------- TTS 音频回调(转发前端) ---------------------- async def tts_audio_callback(audio_chunk, sample_rate, is_finished): """将音频块转发给前端的逻辑""" print(f"音频块长度:{len(audio_chunk)},采样率:{sample_rate},是否结束:{is_finished}") #todo # ---------------------- 定义 ASR 结果回调函数 ---------------------- async def asr_result_callback(result: dict): """ ASR 结果回调:同时转发给前端 + 推送大模型 :param result: ASR 识别结果字典 """ try: print('结果回调:同时转发给前端 + 推送大模型', result) final_asr_text = result["text"] print(f"ASR最终识别结果: {final_asr_text}") # 调用大模型(异步,不阻塞) asyncio.create_task( call_llm_and_send( query=final_asr_text, conversation=llm_conversation ) ) # # 1. 转发给前端(确保连接未断开) # if not websocket.client_state.disconnected: # await websocket.send_json(result) # # # 2. 推送到大模型(过滤错误/空结果) # if not result.get("error") and result.get("text"): # # 异步推送,不阻塞回调 # asyncio.create_task(send_to_llm(result)) except Exception as e: print(f"回调函数执行失败:{e}") # ------------------------------------------------------------------- # ---------------------- 大模型流式回调(转发前端) ---------------------- async def llm_stream_callback(chunk: str, conversation_id: str, is_finished: bool): req_id1 = client.add_tts_request(tts_text=chunk) print('req_id1', req_id1) """ 大模型流式回调:将实时回复转发给前端 """ # if websocket.client_state.disconnected: # return # try: # print('大模型流式回调', chunk) # await websocket.send_json({ # "type": "llm_result", # "data": { # "chunk": chunk, # "conversation_id": conversation_id, # "is_finished": is_finished, # "query": final_asr_text # 关联ASR查询文本 # } # }) # except Exception as e: # print(f"LLM回调执行失败: {e}") # if websocket.client_state.disconnected: # return try: # 1. 转发大模型文本给前端 # await websocket.send_json({ # "type": "llm_result", # "data": { # "chunk": chunk, # "conversation_id": conversation_id, # "is_finished": is_finished, # "user_id": user_id # } # }) # print('大模型返回的', chunk) # await tts_client.send_text_chunk(tts_session_id, chunk) pass # 2. 流式文本发送给 TTS(非错误文本 + 非结束标识) # if chunk and not chunk.startswith("[错误]") and not is_finished: # await tts_client.send_text_chunk(tts_session_id, chunk) # # 合成结束标识 # elif is_finished: # # 发送空文本块触发 TTS 结束 # await tts_client.send_text_chunk(tts_session_id, "") except Exception as e: print(f"LLM回调联动TTS失败({user_id}): {e}") # ---------------------- 调用大模型并转发结果 ---------------------- async def call_llm_and_send(query: str, conversation: LLMConversation): """调用大模型,流式结果通过WebSocket返回前端""" if not query: return print(f"调用大模型 - 用户({user_id}): {query}") # 调用大模型(流式) conv_id, full_reply = await llm_client.send_message( query=query, conversation=conversation, stream_callback=llm_stream_callback, response_mode="streaming" ) print(f"大模型回复完成 - 会话ID: {conv_id}, 完整回复: {full_reply}") # ------------------------------------------------------------------- try: # 1. 获取 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) ) # 3. 并行处理:接收前端数据 + 发送 ASR 结果 async def recv_frontend_data(): """接收前端所有数据,仅将音频数据插入 ASR 队列""" while not asr_conn.stop_event.is_set(): try: raw_bytes = await websocket.receive_bytes() success = await push_audio_data(asr_conn, raw_bytes) if not success: print("音频数据插入失败(队列满/连接失效)") # # 接收前端数据(自动区分类型) # data_type = websocket.receive() # if data_type.type == "bytes": # # 二进制数据 = 音频数据,调用 ASR 接口插入 # raw_bytes = await websocket.receive_bytes() # # 调用 ASR 模块的插入接口,无需管理队列 # success = await push_audio_data(asr_conn, raw_bytes) # if not success: # print("音频数据插入失败(队列满/连接失效)") # elif data_type.type == "json": # # JSON 数据 = 控制指令 # control_data = await websocket.receive_json() # print(f"收到前端控制指令:{control_data}") # if control_data.get("action") == "stop": # asr_conn.stop_event.set() # 触发 ASR 停止 # await websocket.send_json({"text": "", "error": "用户主动停止"}) # break # elif data_type.type == "text": # # 文本数据 = 其他指令 # text_data = await websocket.receive_text() # print(f"收到前端文本数据:{text_data}") except WebSocketDisconnect: print("前端主动断开连接") asr_conn.stop_event.set() break except Exception as e: print(f"接收前端数据失败:{e}") asr_conn.stop_event.set() await websocket.send_json({"error": f"接收数据失败:{str(e)}", "text": ""}) break async def send_asr_result(): """从结果队列取数据,发送给前端""" while not asr_conn.stop_event.is_set(): try: result = await asyncio.wait_for(result_queue.get(), timeout=1.0) # if not websocket.client_state.disconnected: print('result', result) await websocket.send_json(result) except asyncio.TimeoutError: continue except Exception as e: print(f"发送 ASR 结果失败:{e}") asr_conn.stop_event.set() break # 4. 并发执行 await asyncio.gather(recv_frontend_data(), send_asr_result()) except Exception as e: print(f"WebSocket 处理异常:{e}") if asr_conn: asr_conn.stop_event.set() await websocket.send_json({"error": str(e), "text": ""}) finally: # 清理资源 if asr_conn: asr_conn.stop_event.set() # 取消 ASR 任务 if asr_task and not asr_task.done(): asr_task.cancel() try: await asr_task except asyncio.CancelledError: pass # 移除活跃连接 if websocket in active_connections: active_connections.remove(websocket) print(f"客户端连接已清理,当前连接数:{len(active_connections)}") # 关闭前端连接 # if not websocket.client_state.disconnected: # await websocket.close() # try: # # # 4. 持续接收前端消息(二进制,兼容音频/文本/指令) # while not session.is_closed: # raw_bytes = await websocket.receive_bytes() # # print('raw_bytes', raw_bytes) # # todo 前段暂时传入的全部都是语音消息 # if not audio_queue.full(): # audio_queue.put_nowait(raw_bytes) # # await forward_frontend_to_asr(session, raw_bytes) # # 解析前端消息类型(复用之前的编解码器) # # msg_type, data = WsMsgCodec.decode_client_msg(raw_bytes) # # # 分发处理(不同类型消息走不同逻辑) # # if msg_type == ClientMsgType.AUDIO: # # # 音频直接转发给ASR(Session 传递ASR客户端,无用户标识) # # if session.asr_client and not session.is_closed: # # await send_audio_to_asr(session.asr_client, data) # # elif msg_type == ClientMsgType.TEXT: # # # 文本直接走LLM(无需ASR) # # await forward_text_to_llm(session, data) # # elif msg_type == ClientMsgType.ACTION: # # # 操作指令处理 # # await handle_frontend_action(session, data) # # except WebSocketDisconnect: # logger.info(f"会话 [{session.session_id}] 前端主动断开") # except Exception as e: # logger.error(f"会话 [{session.session_id}] 异常: {str(e)}", exc_info=True) # finally: # # 5. 清理会话(一键关闭所有资源,核心价值) # await close_session(session) # await websocket.close()