This commit is contained in:
田岩
2025-12-03 20:54:23 +08:00
parent f250b21b38
commit 11ce05edcf
972 changed files with 121839 additions and 824 deletions
+4 -3
View File
@@ -6,9 +6,10 @@ LOG_LEVEL=INFO
LOG_FILE=logs/app.log
# ASR服务配置
ASR_SERVICE_URL=http://localhost:5000/asr
ASR_TIMEOUT=30 # 超时时间(秒)
ASR_RETRY_TIMES=2 # 重试次数
ASR_HOST=10.10.10.202
ASR_PORT=10096
ASR_TIMEOUT=30 # 超时时间(秒)
ASR_RETRY_TIMES=2 # 重试次数
# LLM服务配置
LLM_SERVICE_URL=http://localhost:6000/chat
+7
View File
@@ -0,0 +1,7 @@
<component name="ProjectDictionaryState">
<dictionary name="project">
<words>
<w>tymas</w>
</words>
</dictionary>
</component>
+57 -17
View File
@@ -1,20 +1,17 @@
from pydantic_settings import BaseSettings, SettingsConfigDict
from pydantic import Field
from pathlib import Path
from typing import List
ROOT_DIR = Path(__file__).parent.parent.parent
class Settings(BaseSettings):
# 应用配置
APP_PORT: int = Field(default=8000, description="服务端口")
LOG_LEVEL: str = Field(default="INFO", description="日志级别")
LOG_FILE: Path = Field(default=ROOT_DIR / "logs/app.log", description="日志文件路径")
# ASR服务配置
ASR_SERVICE_URL: str = Field(..., description="ASR服务地址")
ASR_TIMEOUT: int = Field(default=30, description="ASR超时时间(秒)")
ASR_RETRY_TIMES: int = Field(default=2, description="ASR重试次数")
# LLM服务配置
LLM_SERVICE_URL: str = Field(..., description="LLM服务地址")
LLM_TIMEOUT: int = Field(default=60, description="LLM超时时间(秒)")
@@ -38,20 +35,62 @@ class Settings(BaseSettings):
# -------------------------- 新增:服务版本配置 --------------------------
# ASR当前使用版本(对应ASR_REGISTRY中的key
ASR_CURRENT_VERSION: str = Field(default="local_v1", description="ASR服务当前版本")
# 本地ASR专属配置(仅local_v1版本使用)
LOCAL_ASR_MODEL_PATH: Path = Field(default=ROOT_DIR / "models/asr/local_model", description="本地ASR模型路径")
# 百度云ASR专属配置(仅baidu_v2版本使用)
BAIDU_ASR_API_KEY: str = Field(default="", description="百度云ASR API Key")
BAIDU_ASR_SECRET_KEY: str = Field(default="", description="百度云ASR Secret Key")
ASR_CURRENT_VERSION: str = Field(default="FunASR", description="ASR服务当前版本")
# ASR 公共配置
ASR_TIMEOUT: int = Field(default=5, description="ASR连接超时时间(秒)")
ASR_RETRY_TIMES: int = Field(default=3, description="ASR公共重试次数")
# ASR服务配置
# FunASR 专属配置
ASR_HOST: str = Field(..., description="FunASR服务地址")
ASR_PORT: int = Field(..., description="FunASR服务端口")
ASR_MODE: str = Field(default="2pass", description="FunASR识别模式")
ASR_CHUNK_SIZE: List[int] = Field(default=[5, 10, 5], description="FunASR分片大小配置")
ASR_CHUNK_INTERVAL: int = Field(default=10, description="FunASR分片间隔(毫秒)")
ASR_USE_ITN: int = Field(default=1, description="FunASR是否启用数字转换(1=启用,0=禁用)")
ASR_HOTWORDS: str = Field(default="", description="FunASR热词列表(逗号分隔)")
ASR_RECONNECT_MAX_TIMES: int = Field(default=3, description="FunASR连接重连最大次数")
ASR_POOL_SIZE: int = Field(default=2, description="FunASR连接池大小")
ASR_AUDIO_QUEUE_SIZE: int = Field(default=10000, description="FunASR音频队列最大长度")
# 音频参数(FunASR要求)
ASR_SAMPLE_RATE: int = Field(default=16000, description="音频采样率(Hz")
ASR_CHANNELS: int = Field(default=1, description="音频声道数(1=单声道)")
ASR_SAMPLE_WIDTH: int = Field(default=2, description="音频采样宽度(字节)")
ASR_FRAME_SIZE: int = Field(default=1024, description="音频帧大小")
# LLM当前使用版本(对应LLM_REGISTRY中的key
LLM_CURRENT_VERSION: str = Field(default="local", description="LLM服务当前版本")
# OpenAI LLM专属配置(仅openai版本使用)
OPENAI_API_KEY: str = Field(default="", description="OpenAI API Key")
OPENAI_BASE_URL: str = Field(default="https://api.openai.com/v1", description="OpenAI接口地址")
# 本地LLM专属配置(仅local版本使用)
LOCAL_LLM_MODEL_PATH: Path = Field(default=ROOT_DIR / "models/llm/local_model", description="本地LLM模型路径")
# -------------------------- Dify API 配置(专属) --------------------------
DIFY_BASE_URL: str = Field(
default="http://10.10.10.202:8088/v1",
description="Dify平台API基础URL(如http://xxx:8088/v1"
)
DIFY_API_KEY: str = Field(
default="app-m7HZNV1aGiheh3wr6wNVHFxX",
description="Dify平台API密钥(在Dify应用设置中获取,格式为app-xxx)"
)
DIFY_TIMEOUT: int = Field(
default=30,
description="Dify API请求超时时间(秒)"
)
DIFY_DEFAULT_SCENE: str = Field(
default="通用聊天场景",
description="Dify默认场景描述(传给inputs.scene_description参数)"
)
DIFY_STREAM_CHUNK_SIZE: int = Field(
default=1024,
description="Dify流式响应读取块大小(字节)"
)
# Chat 服务配置(新增)
CHAT_CURRENT_VERSION: str = "DefaultChat" # 对应工厂类的注册名
CHAT_BASE_URL: str = "http://10.10.10.202/v1"
CHAT_API_KEY: str = "app-m7HZNV1aGiheh3wr6wNVHFxX" # 替换为真实 API-Key
CHAT_TIMEOUT: int = 300 # 流式请求超时时间(秒)
CHAT_RETRY_TIMES: int = 3 # 重试次数
CHAT_POOL_SIZE: int = 5 # 连接池大小
# TTS当前使用版本(对应TTS_REGISTRY中的key
TTS_CURRENT_VERSION: str = Field(default="pyttsx3", description="TTS服务当前版本")
@@ -64,7 +103,8 @@ class Settings(BaseSettings):
model_config = SettingsConfigDict(env_file=".env", env_file_encoding="utf-8")
settings = Settings()
# 确保模型目录存在(本地版本需要)
# settings.LOCAL_ASR_MODEL_PATH.parent.mkdir(parents=True, exist_ok=True)
# settings.LOCAL_LLM_MODEL_PATH.parent.mkdir(parents=True, exist_ok=True)
# settings.LOCAL_LLM_MODEL_PATH.parent.mkdir(parents=True, exist_ok=True)
@@ -0,0 +1,48 @@
# audio_ai_chat/asr/asr_manager.py
from typing import Optional,Tuple
from .base import ASRBase
from .factory import ASRFactory
from audio_ai_chat.config.settings import settings
class ASRManager: # 类名与文件名呼应
"""ASR 管理器:负责实例生命周期、连接池复用、资源管理"""
_instance: Optional[ASRBase] = None # 单例存储
@classmethod
async def initialize(cls) -> Tuple[bool, str]:
try:
# 1. 创建 ASR 实例
cls._instance = ASRFactory.get_asr_client()
# 2. 初始化连接池
pool_init_success = await cls._instance.initialize()
if not pool_init_success:
return False, "ASR 连接池初始化失败"
# 3. 异步调用获取有效连接数(添加 await)
valid_conn_count = await cls._instance.get_valid_connection_count()
if valid_conn_count == 0:
return False, f"有效连接数为 0(配置池大小:{settings.ASR_POOL_SIZE}"
return True, f"初始化成功:有效连接数 {valid_conn_count}"
except Exception as e:
return False, f"初始化失败:{str(e)}"
@classmethod
def get_instance(cls) -> Optional[ASRBase]:
"""获取全局 ASR 实例(业务代码调用)"""
return cls._instance
@classmethod
async def close(cls):
"""关闭 ASR 实例和连接池(FastAPI 关闭时调用)"""
if cls._instance:
await cls._instance.close()
cls._instance = None
print("ASR 管理器:实例和连接池已关闭")
# 未来可扩展的管理功能
@classmethod
def is_healthy(cls) -> bool:
"""检查 ASR 实例健康状态(管理功能扩展)"""
return cls._instance is not None
+38 -13
View File
@@ -1,7 +1,12 @@
# audio_ai_chat/asr/base.py
from abc import ABC, abstractmethod
from typing import Optional, Coroutine
from typing import Optional, Dict, Callable, Awaitable
from audio_ai_chat.config.settings import settings
# 定义回调函数类型(异步函数,接收 ASR 结果字典)
ASRResultCallback = Callable[[Dict], Awaitable[None]]
class ASRBase(ABC):
"""ASR服务统一抽象接口"""
def __init__(self):
@@ -10,16 +15,36 @@ class ASRBase(ABC):
self.retry_times = settings.ASR_RETRY_TIMES
@abstractmethod
async def recognize(
self,
voice_data: bytes,
user_id: Optional[str] = None,
**kwargs # 兼容不同版本的额外参数
) -> str:
"""
语音识别核心方法(所有ASR版本必须实现)
:param voice_data: 语音二进制数据
:param user_id: 用户ID(可选)
:return: 识别后的文本
"""
async def initialize(self) -> bool:
"""初始化ASR服务(如连接池初始化)"""
pass
@abstractmethod
async def get_connection(self) -> Optional[object]:
"""获取ASR连接对象"""
pass
@abstractmethod
async def push_audio(self, conn: object, audio_data: bytes) -> bool:
"""推送音频数据到ASR服务"""
pass
@abstractmethod
async def start_communication(self, conn: object, callback: ASRResultCallback) -> None:
"""启动ASR通信(发送音频+接收结果)"""
pass
@abstractmethod
async def release_connection(self, conn: object) -> None:
"""释放ASR连接"""
pass
@abstractmethod
async def close(self) -> None:
"""关闭ASR服务(释放所有连接)"""
pass
@abstractmethod
async def get_valid_connection_count(self) -> int:
"""获取有效连接数(异步方法,子类必须实现)"""
pass
+21 -21
View File
@@ -1,26 +1,26 @@
# audio_ai_chat/asr/factory.py
from typing import Type
from audio_ai_chat.config.settings import settings
from audio_ai_chat.utils.exceptions import ServiceCallError
from .base import ASRBase
# from .version1 import LocalOfflineASR
# from .version2 import BaiduASR
#
# # 注册所有ASR版本:key=配置中的版本名,value=对应的类
# ASR_REGISTRY: dict[str, Type[ASRBase]] = {
# "local_v1": LocalOfflineASR,
# "baidu_v2": BaiduASR,
# # 新增版本时,只需在这里注册:"新版本名": 新类名
# }
from .fun_asr import FunASR # 对应原asr_client.py的实现类
# class ASRFactory:
# """ASR服务工厂类:根据配置创建对应版本的实例"""
# @staticmethod
# def get_asr_client() -> ASRBase:
# # 从配置中获取当前指定的ASR版本
# current_version = settings.ASR_CURRENT_VERSION
# if current_version not in ASR_REGISTRY:
# raise ServiceCallError(
# f"不支持的ASR版本:{current_version},可选版本:{list(ASR_REGISTRY.keys())}"
# )
# # 创建并返回对应版本的实例
# return ASR_REGISTRY[current_version]()
# 注册所有ASR版本:key=配置中的版本名,value=对应的类
ASR_REGISTRY: dict[str, Type[ASRBase]] = {
"FunASR": FunASR,
# 新增版本时,只需在这里注册:"新版本名": 新类名
}
class ASRFactory:
"""ASR服务工厂类:根据配置创建对应版本的实例"""
@staticmethod
def get_asr_client() -> ASRBase:
# 从配置中获取当前指定的ASR版本
current_version = settings.ASR_CURRENT_VERSION
if current_version not in ASR_REGISTRY:
raise ServiceCallError(
f"不支持的ASR版本:{current_version},可选版本:{list(ASR_REGISTRY.keys())}"
)
# 创建并返回对应版本的实例
return ASR_REGISTRY[current_version]()
@@ -0,0 +1 @@
from .fun_asr import FunASR
@@ -0,0 +1,264 @@
import asyncio
import json
import websockets
from typing import Optional, List, Dict, Callable, Awaitable, Any
from dataclasses import dataclass, field
from ..base import ASRBase, ASRResultCallback
from audio_ai_chat.config.settings import settings
# 从配置读取参数(替换原硬编码配置)
AUDIO_PARAMS = {
"sample_rate": settings.ASR_SAMPLE_RATE,
"channels": settings.ASR_CHANNELS,
"sample_width": settings.ASR_SAMPLE_WIDTH,
"frame_size": settings.ASR_FRAME_SIZE
}
ASR_CONFIG = {
"host": settings.ASR_HOST,
"port": settings.ASR_PORT,
"mode": settings.ASR_MODE,
"chunk_size": settings.ASR_CHUNK_SIZE,
"chunk_interval": settings.ASR_CHUNK_INTERVAL,
"use_itn": settings.ASR_USE_ITN,
"hotwords": settings.ASR_HOTWORDS,
"reconnect_max_times": settings.ASR_RECONNECT_MAX_TIMES,
"pool_size": settings.ASR_POOL_SIZE,
"audio_queue_size": settings.ASR_AUDIO_QUEUE_SIZE
}
@dataclass
class ASRConnection:
"""ASR 连接对象(内置音频队列)"""
ws: Optional[websockets.WebSocketClientProtocol] = None
is_busy: bool = False
is_alive: bool = False
reconnect_count: int = 0
audio_queue: asyncio.Queue = field(default_factory=lambda: asyncio.Queue(maxsize=ASR_CONFIG["audio_queue_size"]))
stop_event: asyncio.Event = field(default_factory=asyncio.Event)
class FunASR(ASRBase):
"""FunASR实现类"""
def __init__(self):
super().__init__()
self._connection_pool: List[ASRConnection] = []
self._pool_lock = asyncio.Lock() # 异步锁
async def initialize(self) -> bool:
"""初始化ASR连接池"""
print(f"开始初始化 FunASR 连接池,大小:{ASR_CONFIG['pool_size']}")
tasks = [self._create_single_connection() for _ in range(ASR_CONFIG["pool_size"])]
connections = await asyncio.gather(*tasks)
self._connection_pool = [conn for conn in connections if conn.is_alive]
print(f"FunASR 连接池初始化完成,有效连接数:{len(self._connection_pool)}")
return len(self._connection_pool) > 0
async def get_connection(self) -> Optional[ASRConnection]:
"""从连接池获取空闲连接(实现抽象方法)"""
async with self._pool_lock:
# 查找空闲连接
idle_conns = [
conn for conn in self._connection_pool
if not conn.is_busy and conn.is_alive
]
if idle_conns:
conn = idle_conns[0]
conn.is_busy = True
conn.stop_event.clear()
return conn
# 连接池未满时创建新连接
if len(self._connection_pool) < ASR_CONFIG["pool_size"]:
new_conn = await self._create_single_connection()
if new_conn.is_alive:
new_conn.is_busy = True
self._connection_pool.append(new_conn)
return new_conn
print("FunASR 连接池无空闲连接")
return None
async def push_audio(self, conn: ASRConnection, audio_data: bytes) -> bool:
"""推送音频数据到ASR连接队列(实现抽象方法)"""
if not isinstance(conn, ASRConnection):
print("无效的ASR连接对象")
return False
if not conn or not conn.is_alive or conn.stop_event.is_set():
return False
try:
conn.audio_queue.put_nowait(audio_data)
return True
except asyncio.QueueFull:
print("FunASR 音频队列已满,丢弃当前音频帧")
return False
async def start_communication(self, conn: ASRConnection, callback: ASRResultCallback) -> None:
"""启动ASR通信(发送音频+接收结果,实现抽象方法)"""
if not isinstance(conn, ASRConnection):
await callback({"error": "无效的ASR连接对象", "text": ""})
return
await self._handle_communication(conn, callback)
async def release_connection(self, conn: ASRConnection) -> None:
"""释放ASR连接(实现抽象方法)"""
if not isinstance(conn, ASRConnection):
print("无效的ASR连接对象,无法释放")
return
async with self._pool_lock:
conn.is_busy = False
conn.stop_event.set()
# 清空队列
while not conn.audio_queue.empty():
try:
conn.audio_queue.get_nowait()
except asyncio.QueueEmpty:
break
# 重连逻辑
if not conn.is_alive and conn.reconnect_count < ASR_CONFIG["reconnect_max_times"]:
print(f"尝试重连 FunASR 连接(次数:{conn.reconnect_count + 1}")
new_conn = await self._create_single_connection()
if new_conn.is_alive:
if conn in self._connection_pool:
idx = self._connection_pool.index(conn)
self._connection_pool[idx] = new_conn
else:
conn.reconnect_count += 1
elif conn.reconnect_count >= ASR_CONFIG["reconnect_max_times"]:
if conn in self._connection_pool:
self._connection_pool.remove(conn)
print("FunASR 连接重连次数耗尽,已移除")
async def close(self) -> None:
"""关闭所有ASR连接(实现抽象方法)"""
async with self._pool_lock:
# for conn in self._connection_pool:
# conn.stop_event.set()
# if conn.ws and not conn.ws.closed:
# try:
# await conn.ws.close()
# print("FunASR 连接已关闭")
# except Exception as e:
# print(f"关闭 FunASR 连接失败:{e}")
self._connection_pool.clear()
print("FunASR 连接池已清空")
async def _create_single_connection(self) -> ASRConnection:
"""创建单个ASR连接(内部私有方法)"""
asr_conn = ASRConnection()
asr_uri = f"ws://{ASR_CONFIG['host']}:{ASR_CONFIG['port']}"
try:
ws = await websockets.connect(
asr_uri,
subprotocols=["binary"],
ping_interval=None,
open_timeout=self.timeout # 使用基类的超时配置
)
asr_conn.ws = ws
asr_conn.is_alive = True
# 发送初始化配置
init_msg = json.dumps({
"mode": ASR_CONFIG["mode"],
"chunk_size": ASR_CONFIG["chunk_size"],
"chunk_interval": ASR_CONFIG["chunk_interval"],
"wav_name": "pool_connection",
"is_speaking": True,
"hotwords": ASR_CONFIG["hotwords"],
"itn": bool(ASR_CONFIG["use_itn"]),
"audio_fs": AUDIO_PARAMS["sample_rate"]
})
await ws.send(init_msg)
print("FunASR 连接初始化成功")
return asr_conn
except Exception as e:
print(f"创建 FunASR 连接失败:{e}")
asr_conn.is_alive = False
return asr_conn
async def _handle_communication(
self,
asr_conn: ASRConnection,
result_callback: ASRResultCallback
):
"""处理ASR通信细节(内部私有方法)"""
if not asr_conn or not asr_conn.ws:
await result_callback({"error": "无可用 FunASR 连接", "text": ""})
return
# 发送音频任务
async def send_audio():
while not asr_conn.stop_event.is_set() and asr_conn.is_alive:
try:
pcm_data = await asyncio.wait_for(asr_conn.audio_queue.get(), timeout=1.0)
if pcm_data and asr_conn.is_alive:
await asr_conn.ws.send(pcm_data)
await asyncio.sleep(0.005)
except asyncio.TimeoutError:
continue
except Exception as e:
print(f"发送音频到 FunASR 失败:{e}")
asr_conn.is_alive = False
await result_callback({"error": f"音频发送失败:{str(e)}", "text": ""})
asr_conn.stop_event.set()
break
# 接收结果任务
async def recv_result():
while not asr_conn.stop_event.is_set() and asr_conn.is_alive:
try:
asr_result = await asr_conn.ws.recv()
result_json = json.loads(asr_result)
# 打印asr实时结果
# print(result_json.get("text", ""))
if result_json.get("timestamp", "") == '':
continue
result = {
"text": result_json.get("text", ""),
"mode": result_json.get("mode", ""),
"timestamp": result_json.get("timestamp", ""),
"is_final": result_json.get("is_final", True),
"error": ""
}
await result_callback(result)
except websockets.exceptions.ConnectionClosed:
print("FunASR 连接已关闭")
asr_conn.is_alive = False
await result_callback({"error": "FunASR 连接断开", "text": ""})
asr_conn.stop_event.set()
break
except Exception as e:
print(f"接收 FunASR 结果失败:{e}")
asr_conn.is_alive = False
await result_callback({"error": f"接收结果失败:{str(e)}", "text": ""})
asr_conn.stop_event.set()
break
try:
send_task = asyncio.create_task(send_audio())
recv_task = asyncio.create_task(recv_result())
await asyncio.gather(send_task, recv_task)
finally:
# 确保任务被取消
send_task.cancel()
recv_task.cancel()
try:
await send_task
await recv_task
except asyncio.CancelledError:
pass
# 自动释放连接
await self.release_connection(asr_conn)
async def get_valid_connection_count(self) -> int:
"""获取有效连接数(异步方法,保证线程安全)"""
async with self._pool_lock: # 异步锁,自动 acquire/release
# 过滤出 "存活" 且 "在连接池内" 的连接
valid_conns = [conn for conn in self._connection_pool if conn.is_alive]
return len(valid_conns)
+360 -51
View File
@@ -1,9 +1,21 @@
from typing import Optional, Dict, Any, List
from typing import Optional, Dict, Any, List, TypedDict
import asyncio
from datetime import datetime
from audio_ai_chat.config.logger import logger
from audio_ai_chat.core.llm.dify.dify import LLMConversation
# from audio_ai_chat.core.llm.factory import LLMFactory
from audio_ai_chat.core.llm.base import LLMBase
# 定义对话历史条目类型(TypedDict 用于类型提示,更清晰)
class ChatHistoryItem(TypedDict):
"""对话历史条目结构(强类型定义)"""
role: str # 发言人角色:"user"(用户)、"assistant"(助手)、"system"(系统)
content: str # 对话内容(ASR转写结果/大模型回复/系统提示)
timestamp: str # 对话时间(ISO 8601格式,如 "2024-05-20T14:30:00.123Z"
source: str # 内容来源:"asr"(语音转写)、"text"(纯文本输入)、"llm"(大模型生成)、"system"(系统配置)
def get_current_iso_timestamp() -> str:
"""获取当前时间的ISO 8601格式字符串(UTC时间)"""
return datetime.utcnow().isoformat(timespec="milliseconds") + "Z"
class ConnectionContext:
"""
@@ -15,30 +27,43 @@ class ConnectionContext:
"""
初始化连接上下文
:param client_id: WebSocket连接唯一标识(如id(websocket)
:param user_id: 用户唯一标识(从前端请求中获取)
"""
self.MAX_CHAT_HISTORY = 100 # 单个连接最大对话历史条数
self.MAX_QUEUE_SIZE = 50 # 单个连接消息队列最大长度
self.client_id = client_id # 连接唯一ID
self.created_at = asyncio.get_event_loop().time() # 连接创建时间
self.created_at = asyncio.get_event_loop().time() # 连接创建时间(时间戳)
self.created_at_str = datetime.utcnow().isoformat() + "Z" # 连接创建时间(ISO格式)
# 1. 大模型独立Session(每个连接创建一个新的LLM客户端实例)
# self.llm_session: LLMBase = LLMFactory.get_llm_client() # 独立Session
self.chat_history: List[Dict[str, str]] = [] # 该连接的对话历史([(user: "...", assistant: "..."), ...]
self.llm_session: Optional[LLMConversation] = None # 实际是DifyLLMClient实例
# 优化后的对话历史:List[ChatHistoryItem]
self.chat_history: List[ChatHistoryItem] = [] # 该连接的完整对话历史
# 2. 异步消息队列(用于缓存TTS结果,有序推送给前端)
self.tts_client = None
self.message_queue: asyncio.Queue[bytes] = asyncio.Queue()
# 3. 连接状态(可选:如是否正在处理请求、是否断开等)
# 3. 连接状态
self.is_active: bool = True
self.is_processing: bool = False
self.name = None
self.user_id = None
self.token = None
self.name: Optional[str] = None # 用户名
self.user_id: Optional[str] = None # 用户唯一标识
self.token: Optional[str] = None # 用户令牌
self.disconnect_time: Optional[datetime] = None # 断开时间(None 表示活跃)
# 4. ASR临时缓存(处理流式结果,避免重复存储)
self._current_asr_text: str = "" # 当前正在拼接的ASR文本
self._current_asr_metadata: Optional[Dict[str, Any]] = None # 当前ASR元数据
self.is_processing: bool = False
def set_user_info(self, token: str, user_id: str, name: str = "匿名用户"):
"""
二次设置用户信息(身份校验通过后调用)
:param token:
:param token: 用户令牌
:param user_id: 用户唯一标识(必填)
:param name: 用户名(可选,默认匿名)
"""
@@ -49,6 +74,28 @@ class ConnectionContext:
self.token = token
logger.debug(f"客户端 {self.client_id} 设置用户信息:user_id={user_id}, name={name}")
def init_llm_session(self):
"""初始化Dify客户端(每个连接一个实例,存入context)"""
# if not self.user_id:
# raise InitError(f"客户端 {self.client_id} 未设置用户信息,无法初始化Dify客户端")
if self.llm_session:
logger.warning(f"客户端 {self.client_id} Dify客户端已存在,无需重复初始化")
return
# 工厂类创建Dify客户端实例,存入当前context
self.llm_session = LLMFactory.get_llm_client(version="dify")
logger.debug(f"客户端 {self.client_id}user_id={self.user_id}Dify客户端初始化完成")
def update_dify_context(self, user_text: str, assistant_text: str, conversation_id: Optional[str]):
"""更新Dify会话上下文(隔离存储)"""
if conversation_id:
self.dify_conversation_id = conversation_id
# 限制历史长度(最多50轮)
self.chat_history.append({"user": user_text, "assistant": assistant_text})
if len(self.chat_history) > 50:
self.chat_history.pop(0)
logger.debug(
f"客户端 {self.client_id} Dify上下文更新:conversation_id={self.dify_conversation_id},历史长度={len(self.dify_history)}")
async def add_message_to_queue(self, message: bytes):
"""将TTS结果添加到消息队列(异步安全)"""
if not self.is_active:
@@ -56,6 +103,15 @@ class ConnectionContext:
await self.message_queue.put(message)
logger.debug(f"消息队列添加数据:client_id={self.client_id},队列长度={self.message_queue.qsize()}")
def complete_initialization(self):
"""标记完成初始化(必须确保Dify客户端已创建)"""
# if not self.user_id:
# raise InitError(f"客户端 {self.client_id} 未设置用户信息")
# if not self.llm_session:
# raise InitError(f"客户端 {self.client_id} 未初始化Dify客户端")
self.is_initialized = True
logger.info(f"客户端 {self.client_id}user_id={self.user_id})完整初始化完成")
async def get_message_from_queue(self) -> Optional[bytes]:
"""从消息队列获取消息(异步阻塞,直到有消息或连接断开)"""
try:
@@ -65,81 +121,334 @@ class ConnectionContext:
logger.debug(f"消息队列超时:client_id={self.client_id},无新消息")
return None
def update_chat_history(self, user_text: str, assistant_text: str):
"""更新该连接的对话历史"""
self.chat_history.append({
"user": user_text,
"assistant": assistant_text
})
# 可选:限制历史长度(避免内存溢出)
if len(self.chat_history) > 50:
self.chat_history.pop(0) # 删除最早的历史
def add_chat_history(self, item: ChatHistoryItem):
"""
添加对话历史条目(统一接口,支持用户/助手/系统消息)
:param item: 符合 ChatHistoryItem 结构的对话条目
"""
# 补全必填字段(防止遗漏)
if "timestamp" not in item:
item["timestamp"] = datetime.utcnow().isoformat() + "Z"
if "source" not in item:
item["source"] = "unknown"
self.chat_history.append(item)
# 可选:限制历史长度(避免内存溢出,保留最近100条)
if len(self.chat_history) > 100:
removed_item = self.chat_history.pop(0)
logger.debug(
f"对话历史超出限制,删除最早条目:{removed_item['timestamp']} - {removed_item['role']}: {removed_item['content'][:20]}...")
logger.debug(
f"添加对话历史:client_id={self.client_id}"
f"role={item['role']}content={item['content'][:30]}..."
)
def add_asr_result(self, asr_result: Dict[str, Any]):
"""
处理ASR结果,拼接流式文本,最终结果存入对话历史
:param asr_result: ASR返回的结果字典(含text、is_final、timestamp等)
"""
if asr_result.get("error"):
logger.error(f"ASR错误:client_id={self.client_id}error={asr_result['error']}")
return
# 提取ASR核心信息
asr_text = asr_result.get("text", "").strip()
is_final = asr_result.get("is_final", False)
asr_timestamp = asr_result.get("timestamp", "")
asr_mode = asr_result.get("mode", "")
# 缓存ASR元数据(流式过程中更新)
self._current_asr_metadata = {
"timestamp": asr_timestamp,
"mode": asr_mode,
"is_final": is_final,
"source": "asr"
}
# 拼接流式文本(处理部分结果)
if asr_text:
# 避免重复拼接(如果ASR返回重复文本)
if not self._current_asr_text.endswith(asr_text) and self._current_asr_text != asr_text:
self._current_asr_text += asr_text if not self._current_asr_text else f" {asr_text}"
# 当ASR返回最终结果时,存入对话历史
if is_final:
if self._current_asr_text:
# 构造对话历史条目
chat_item: ChatHistoryItem = {
"role": "user", # ASR结果属于用户输入
"content": self._current_asr_text,
"timestamp": datetime.utcnow().isoformat() + "Z",
"source": "asr",
"asr_metadata": self._current_asr_metadata
}
# 添加到对话历史
self.add_chat_history(chat_item)
# 清空临时缓存
self._current_asr_text = ""
self._current_asr_metadata = None
else:
logger.warning(f"ASR最终结果为空:client_id={self.client_id}")
def add_llm_result(self, llm_text: str):
"""
添加大模型回复到对话历史
:param llm_text: 大模型生成的回复文本
"""
if not llm_text.strip():
logger.warning(f"大模型回复为空:client_id={self.client_id}")
return
chat_item: ChatHistoryItem = {
"role": "assistant", # 大模型回复属于助手角色
"content": llm_text.strip(),
"timestamp": datetime.utcnow().isoformat() + "Z",
"source": "llm",
"asr_metadata": None # 大模型回复无ASR元数据
}
self.add_chat_history(chat_item)
def add_system_message(self, system_text: str):
"""
添加系统消息到对话历史(如错误提示、系统通知)
:param system_text: 系统消息文本
"""
chat_item: ChatHistoryItem = {
"role": "system", # 系统角色
"content": system_text.strip(),
"timestamp": datetime.utcnow().isoformat() + "Z",
"source": "system",
"asr_metadata": None
}
self.add_chat_history(chat_item)
def get_chat_history(self, limit: Optional[int] = None) -> List[ChatHistoryItem]:
"""
获取对话历史(支持限制返回条数)
:param limit: 限制返回的最新条数,None表示返回全部
:return: 过滤后的对话历史
"""
if limit and isinstance(limit, int) and limit > 0:
return self.chat_history[-limit:] # 返回最近N条
return self.chat_history.copy() # 返回全部(拷贝,避免外部修改)
def close(self):
"""关闭连接上下文,释放资源"""
self.is_active = False
self.is_processing = False
# 清空消息队列(可选)
# 清空消息队列
while not self.message_queue.empty():
try:
self.message_queue.get_nowait()
except asyncio.QueueEmpty:
break
logger.info(f"连接上下文已关闭:client_id={self.client_id}user_id={self.user_id}")
# 记录连接关闭日志(包含对话历史统计)
logger.info(
f"连接上下文已关闭:client_id={self.client_id}"
f"user_id={self.user_id}"
f"对话历史条数={len(self.chat_history)}"
)
def mark_disconnected(self):
self.is_active = False
self.disconnect_time = datetime.utcnow()
logger.info(f"客户端 {self.client_id} 标记为断开,待延迟清理(user_id={self.user_id}")
def __del__(self):
"""析构函数:确保资源释放"""
self.close()
# -------------------------- 核心:调用Dify流式接口 --------------------------
async def call_dify_stream(
self,
user_text: str,
tts_client,
asr_metadata: Optional[Dict[str, Any]] = None # 接收ASR元数据
) -> str:
"""
调用Dify流式接口,使用ChatHistoryItem存储完整历史
:param user_text: ASR识别后的用户文本
:param tts_client: TTS客户端实例
:param asr_metadata: ASR元数据(如置信度、语音时长等)
:return: Dify完整回复文本
"""
if not self.is_initialized:
raise InitError(f"客户端 {self.client_id} 未完成初始化,无法调用Dify")
if not user_text:
raise ValueError("用户输入文本不能为空")
full_response = ""
llm_metadata: Dict[str, Any] = {} # 存储Dify元数据
# 1. 添加用户输入到对话历史(user角色,source=asr
user_history_item: ChatHistoryItem = {
"role": "user",
"content": user_text,
"timestamp": get_current_iso_timestamp(),
"source": "asr",
"asr_metadata": asr_metadata, # 传入ASR元数据
"llm_metadata": None
}
self.add_chat_history_item(user_history_item)
# -------------------------- 流式回调函数(闭包访问context --------------------------
async def stream_callback(chunk: str, conversation_id: str, is_finished: bool):
nonlocal full_response, llm_metadata
if chunk and not is_finished:
# 累加完整回复
full_response += chunk
logger.debug(
f"客户端 {self.client_id} Dify流式片段:content_len={len(chunk)}, "
f"累计_len={len(full_response)}"
)
# 实时调用TTS合成音频
try:
tts_audio = await tts_client.synthesize(
text=chunk,
user_id=self.user_id
)
await self.add_tts_to_queue(tts_audio)
except Exception as e:
logger.error(f"客户端 {self.client_id} TTS合成失败:{str(e)}")
return
# 流式结束:添加助手回复到对话历史
if is_finished and full_response:
# 更新Dify会话ID和元数据
self.dify_conversation_id = conversation_id
llm_metadata = {
"conversation_id": conversation_id,
"response_mode": "streaming",
"full_response_len": len(full_response),
"timestamp": get_current_iso_timestamp()
}
# 添加助手回复到对话历史(assistant角色,source=llm
assistant_history_item: ChatHistoryItem = {
"role": "assistant",
"content": full_response,
"timestamp": get_current_iso_timestamp(),
"source": "llm",
"asr_metadata": None,
"llm_metadata": llm_metadata # 存储Dify元数据
}
self.add_chat_history_item(assistant_history_item)
# -------------------------- 调用Dify流式接口 --------------------------
try:
await self.llm_session.chat(
text=user_text,
user_id=self.user_id,
history=self.get_dify_compatible_history(), # 传入Dify兼容格式的历史
stream_callback=stream_callback,
response_mode="streaming"
)
except Exception as e:
logger.error(f"客户端 {self.client_id} Dify流式调用失败:{str(e)}")
raise
return full_response
# -------------------------- 关键修改:ConnectionManager 全局单例 --------------------------
class ConnectionManager:
"""
WebSocket连接全局管理器:维护所有活跃连接的上下文
提供创建、查询、删除连接上下文的接口(线程/异步安全)
"""
"""全局连接上下文管理器(支持延迟清理和重连复用)"""
_instance: Optional["ConnectionManager"] = None
_lock = asyncio.Lock() # 单例锁
def __new__(cls):
raise NotImplementedError("请使用 ConnectionManager.get_instance() 获取实例")
def __init__(self):
# 存储所有活跃连接key=client_idint),value=ConnectionContext实例
self.connections: Dict[int, ConnectionContext] = {}
# 异步锁:确保多连接并发操作时的数据安全
self._lock = asyncio.Lock()
# 存储所有上下文key=client_id当前活跃连接的唯一标识)
self.active_contexts: Dict[str, ConnectionContext] = {}
# 存储待清理的上下文:key=user_id(用户唯一标识,用于重连匹配)
self.pending_clean_contexts: Dict[str, ConnectionContext] = {}
self._internal_lock = asyncio.Lock() # 操作锁
async def create_connection(self, client_id: int, user_id: Optional[str] = None) -> ConnectionContext:
"""创建新的连接上下文(线程安全)"""
async with self._lock:
# 避免重复创建(同一client_id不会重复连接)
@classmethod
async def get_instance(cls) -> "ConnectionManager":
"""获取全局唯一实例(异步安全)"""
if cls._instance is None:
async with cls._lock:
if cls._instance is None: # 双重检查锁定
cls._instance = super().__new__(cls)
cls._instance.__init__()
logger.info("ConnectionManager 全局单例初始化成功")
return cls._instance
async def create_connection(self, client_id: str) -> ConnectionContext:
"""创建连接上下文(异步安全)"""
async with self._internal_lock:
if client_id in self.connections:
logger.warning(f"连接已存在:client_id={client_id},将覆盖旧连接")
self.connections[client_id].close()
# 创建新的连接上下文(包含独立LLM Session和消息队列)
context = ConnectionContext(client_id=client_id, user_id=user_id)
context = ConnectionContext(client_id=client_id)
self.connections[client_id] = context
logger.info(
f"创建新连接上下文:client_id={client_id}user_id={user_id},当前活跃连接数={len(self.connections)}")
logger.info(f"创建连接上下文:client_id={client_id},活跃连接数={len(self.connections)}")
return context
async def get_connection(self, client_id: int) -> Optional[ConnectionContext]:
"""获取指定client_id的连接上下文(线程安全)"""
async with self._lock:
async def get_connection(self, client_id: str) -> Optional[ConnectionContext]:
"""获取连接上下文(异步安全)"""
async with self._internal_lock:
context = self.connections.get(client_id)
if context and not context.is_active:
# 清理已断开的连接
del self.connections[client_id]
return None
return context
async def remove_connection(self, client_id: int):
"""除连接上下文(线程安全)"""
async with self._lock:
async def remove_connection(self, client_id: str):
"""除连接上下文(异步安全)"""
async with self._internal_lock:
context = self.connections.pop(client_id, None)
if context:
context.close()
logger.info(f"移除连接上下文:client_id={client_id}当前活跃连接数={len(self.connections)}")
logger.info(f"移除连接上下文:client_id={client_id},活跃连接数={len(self.connections)}")
async def get_active_connections_count(self) -> int:
"""获取当前活跃连接数(线程安全)"""
async with self._lock:
# 过滤已断开的连接
"""获取活跃连接数(异步安全)"""
async with self._internal_lock:
self.connections = {k: v for k, v in self.connections.items() if v.is_active}
return len(self.connections)
return len(self.connections)
# 检查是否在可重连时间窗口内(30分钟)
# 标记为断开连接(不立即清理)
def is_reconnectable(self, timeout: int = 30) -> bool:
if self.is_active or not self.disconnect_time:
return False # 活跃连接或未记录断开时间,不可重连
# 计算断开时间是否在 timeout 分钟内
return datetime.utcnow() - self.disconnect_time <= timedelta(minutes=timeout)
async def create_or_reconnect_context(self, new_client_id: str, user_id: Optional[str] = None) -> ConnectionContext:
"""
创建新上下文或重连复用旧上下文
:param new_client_id: 新 WebSocket 连接的 client_id
:param user_id: 用户唯一标识(用于匹配旧上下文)
:return: 新上下文或复用的旧上下文
"""
async with self._internal_lock:
# 1. 如果用户已登录(有 user_id),先尝试重连复用
if user_id and user_id in self.pending_clean_contexts:
old_context = self.pending_clean_contexts.pop(user_id)
if old_context.is_reconnectable():
# 重连激活旧上下文,更新 client_id
old_context.reconnect(new_client_id=new_client_id)
# 加入活跃上下文列表
self.active_contexts[new_client_id] = old_context
return old_context
else:
# 旧上下文已超时,清理并创建新的
old_context.close()
# 2. 无旧上下文可复用,创建新上下文
new_context = ConnectionContext(client_id=new_client_id)
if user_id:
new_context.user_id = user_id # 绑定用户标识(如果已提供)
self.active_contexts[new_client_id] = new_context
logger.info(f"创建新上下文:client_id={new_client_id}user_id={user_id}")
return new_context
+64 -19
View File
@@ -1,26 +1,71 @@
from abc import ABC, abstractmethod
from typing import Optional, Coroutine
from typing import Optional, Dict, Callable, Awaitable
from audio_ai_chat.config.settings import settings
class LLMBase(ABC):
"""LLM服务统一抽象接口"""
# 定义回调函数类型(异步函数,与 ASR 回调风格一致)
# ChatStreamCallback = Callable[[str, bool, ConnectionContext], Awaitable[None]]
"""
Chat 流式回调函数类型:
- 第一个参数:流式文本块
- 第二个参数:是否结束标记
- 第三个参数:上下文对象
"""
class ChatBase(ABC):
"""Chat 服务统一抽象接口(与 ASRBase 接口风格完全对齐)"""
def __init__(self):
self.timeout = settings.LLM_TIMEOUT
self.retry_times = settings.LLM_RETRY_TIMES
self.model = settings.LLM_MODEL # 模型版本(不同LLM可能支持不同模型)
# 公共配置(所有 Chat 实现共享)
self.timeout = settings.CHAT_TIMEOUT # 需在配置中添加 CHAT_TIMEOUT
self.retry_times = settings.CHAT_RETRY_TIMES # 需在配置中添加 CHAT_RETRY_TIMES
self.base_url = settings.CHAT_BASE_URL # 配置中添加:http://10.10.10.202/v1
self.api_key = settings.CHAT_API_KEY # 配置中添加 API-Key
@abstractmethod
async def chat(
async def initialize(self) -> bool:
"""初始化 Chat 服务(如连接池、全局配置)"""
pass
@abstractmethod
async def get_connection(self) -> Optional[object]:
"""获取 Chat 连接对象(与 ASR 的 get_connection 对应)"""
pass
@abstractmethod
async def send_message(
self,
text: str,
user_id: Optional[str] = None,
history: Optional[list] = None, # 对话历史(部分LLM支持)
**kwargs
) -> str:
"""
大模型对话核心方法
:param text: 用户输入文本(ASR识别结果)
:param user_id: 用户ID(可选)
:param history: 对话历史(可选,格式:[(用户输入, 模型回答), ...])
:return: 模型生成的回答文本
"""
conn: object,
query: str,
context: ConnectionContext,
inputs: Optional[Dict] = None
) -> bool:
"""发送聊天消息(类似 ASR 的 push_audio"""
pass
@abstractmethod
async def start_communication(
self,
conn: object,
callback: ChatStreamCallback,
query: str,
context: ConnectionContext,
inputs: Optional[Dict] = None
) -> None:
"""启动 Chat 通信(流式接收结果,与 ASR 的 start_communication 对应)"""
pass
@abstractmethod
async def release_connection(self, conn: object) -> None:
"""释放 Chat 连接(与 ASR 的 release_connection 对应)"""
pass
@abstractmethod
async def close(self) -> None:
"""关闭 Chat 服务(释放所有连接,与 ASR 的 close 对应)"""
pass
@abstractmethod
async def get_valid_connection_count(self) -> int:
"""获取有效连接数(与 ASR 的接口完全一致)"""
pass
@@ -0,0 +1 @@
# from .dify import DifyLLMClient
@@ -0,0 +1,175 @@
import asyncio
import json
from typing import Optional, Dict, Callable, Awaitable
from dataclasses import dataclass, field
import aiohttp # 新增:异步HTTP库
# 大模型配置(集中管理)
LLM_CONFIG = {
"base_url": "http://10.10.10.202:8088/v1",
"api_key": "app-m7HZNV1aGiheh3wr6wNVHFxX",
"timeout": 30, # 请求超时时间(秒)
"default_scene": "通用聊天场景", # 默认场景描述
"stream_chunk_size": 1024 # 流式接收块大小
}
# 定义流式回调函数类型(异步)
LLMStreamCallback = Callable[[str, Optional[str], bool], Awaitable[None]]
"""
回调函数参数说明:
- chunk: 单次流式返回的文本片段
- conversation_id: 会话ID(首次返回,后续复用)
- is_finished: 是否结束(True=流式结束/同步返回完成)
"""
@dataclass
class LLMConversation:
"""会话对象(管理会话ID和上下文)"""
conversation_id: Optional[str] = None
user_id: str = ""
scene_description: str = LLM_CONFIG["default_scene"]
# 可选:存储会话历史(如需上下文管理)
history: list = field(default_factory=list)
class LLMClient:
"""大模型客户端封装(异步修复版)"""
def __init__(self):
self.base_url = LLM_CONFIG["base_url"]
self.headers = {
"Authorization": f"Bearer {LLM_CONFIG['api_key']}",
"Content-Type": "application/json"
}
self.timeout = aiohttp.ClientTimeout(total=LLM_CONFIG["timeout"]) # 异步超时
self._session: Optional[aiohttp.ClientSession] = None # 异步会话(复用连接)
async def _get_session(self) -> aiohttp.ClientSession:
"""获取/复用异步HTTP会话"""
if self._session is None or self._session.closed:
self._session = aiohttp.ClientSession(timeout=self.timeout)
return self._session
async def send_message(
self,
query: str,
conversation: LLMConversation,
stream_callback: Optional[LLMStreamCallback] = None,
response_mode: str = "streaming"
) -> tuple[Optional[str], str]:
"""
发送消息到大模型(纯异步版,无线程池阻塞)
"""
payload = {
"query": query,
"inputs": {"scene_description": conversation.scene_description},
"response_mode": response_mode,
"user": conversation.user_id
}
if conversation.conversation_id:
payload["conversation_id"] = conversation.conversation_id
url = f"{self.base_url}/chat-messages"
full_response = ""
res_conversation_id = conversation.conversation_id
try:
session = await self._get_session()
if response_mode == "streaming":
# 异步流式请求(无线程池,纯异步IO)
async with session.post(url, headers=self.headers, json=payload) as response:
response.raise_for_status()
# 实时迭代流式响应
async for line in response.content.iter_chunked(LLM_CONFIG["stream_chunk_size"]):
if not line:
continue
line_data = line.decode("utf-8")
if line_data.startswith("data: "):
json_str = line_data[6:].strip()
if json_str == "[DONE]":
if stream_callback:
await stream_callback("", res_conversation_id, True)
break
try:
data = json.loads(json_str)
# 更新会话ID
if not res_conversation_id and "conversation_id" in data:
res_conversation_id = data["conversation_id"]
# 提取内容
chunk = data.get("content", data.get("answer", data.get("message", "")))
# print('大模型返回的', chunk)
if chunk:
full_response += chunk
if stream_callback:
await stream_callback(chunk, res_conversation_id, False)
await asyncio.sleep(0) # 让出调度权
except json.JSONDecodeError as e:
# print(f"大模型解析流式数据失败: {e}")
continue
else:
# 异步非流式请求
async with session.post(url, headers=self.headers, json=payload) as response:
response.raise_for_status()
data = await response.json()
res_conversation_id = data.get("conversation_id", conversation.conversation_id)
full_response = data.get("content", data.get("answer", data.get("message", "")))
if stream_callback:
await stream_callback(full_response, res_conversation_id, True)
conversation.conversation_id = res_conversation_id
return res_conversation_id, full_response
except aiohttp.ClientError as e:
error_msg = f"大模型请求失败: {str(e)}"
print(error_msg)
if stream_callback:
await stream_callback(f"[错误] {error_msg}", res_conversation_id, True)
return res_conversation_id, ""
except Exception as e:
error_msg = f"大模型处理异常: {str(e)}"
print(error_msg)
if stream_callback:
await stream_callback(f"[错误] {error_msg}", res_conversation_id, True)
return res_conversation_id, ""
async def close(self):
"""关闭异步会话(程序退出时调用)"""
if self._session and not self._session.closed:
await self._session.close()
# 全局单例客户端(异步版)
llm_client = LLMClient()
# 快捷调用函数(保持原有接口不变)
async def call_llm(
query: str,
user_id: str,
scene_description: str = LLM_CONFIG["default_scene"],
conversation_id: Optional[str] = None,
stream_callback: Optional[LLMStreamCallback] = None,
response_mode: str = "streaming"
) -> tuple[Optional[str], str]:
"""
快捷调用大模型(无需手动创建会话对象)
:param query: 用户提问
:param user_id: 用户ID
:param scene_description: 场景描述
:param conversation_id: 会话ID(续聊用)
:param stream_callback: 流式回调
:param response_mode: 响应模式
:return: (conversation_id, 完整回复)
"""
conversation = LLMConversation(
conversation_id=conversation_id,
user_id=user_id,
scene_description=scene_description
)
return await llm_client.send_message(
query=query,
conversation=conversation,
stream_callback=stream_callback,
response_mode=response_mode
)
# 可选:程序退出时关闭会话(如FastAPI的shutdown事件)
async def shutdown_llm_client():
await llm_client.close()
+22 -19
View File
@@ -1,22 +1,25 @@
from typing import Type
from audio_ai_chat.config.settings import settings
from audio_ai_chat.utils.exceptions import ServiceCallError
from .base import LLMBase
#
# from .openai_llm import OpenAILLM
# from .local_llm import LocalLLM
#
# LLM_REGISTRY: dict[str, Type[LLMBase]] = {
# "openai": OpenAILLM,
# "local": LocalLLM,
# }
#
# class LLMFactory:
# @staticmethod
# def get_llm_client() -> LLMBase:
# current_version = settings.LLM_CURRENT_VERSION
# if current_version not in LLM_REGISTRY:
# raise ServiceCallError(
# f"不支持的LLM版本:{current_version},可选版本:{list(LLM_REGISTRY.keys())}"
# )
# return LLM_REGISTRY[current_version]()
from .base import ChatBase
from .dify import DifyLLMClient # 具体实现类(对应 ASR 的 FunASR)
# 注册所有 Chat 实现:key=配置中的版本名,value=对应的类
CHAT_REGISTRY: dict[str, Type[ChatBase]] = {
"DefaultChat": DifyLLMClient,
# 新增 Chat 实现时,只需在这里注册
}
class ChatFactory:
"""Chat 服务工厂类(与 ASRFactory 逻辑完全一致)"""
@staticmethod
def get_chat_client() -> ChatBase:
# 从配置中获取当前指定的 Chat 版本
current_version = settings.CHAT_CURRENT_VERSION # 配置中添加该字段
if current_version not in CHAT_REGISTRY:
raise ServiceCallError(
f"不支持的 Chat 版本:{current_version},可选版本:{list(CHAT_REGISTRY.keys())}"
)
# 创建并返回对应版本的实例
return CHAT_REGISTRY[current_version]()
@@ -0,0 +1,75 @@
from typing import Optional, Tuple, Dict
from .base import ChatBase
from .factory import ChatFactory
from audio_ai_chat.config.settings import settings
# # 上下文管理类(集成到管理器中,与业务逻辑解耦)
# @dataclass
# class ConnectionContext:
# """Chat 连接上下文管理类(与你的原有 Context 兼容)"""
# user_id: str
# conversation_id: Optional[str] = None
# current_context: Dict = None
# # 可添加更多业务字段(如请求ID、会话状态等)
#
# def __post_init__(self):
# if self.current_context is None:
# self.current_context = {}
#
# def update_conversation_id(self, conversation_id: str):
# """更新会话ID(历史对话用)"""
# self.conversation_id = conversation_id
# self.current_context['conversation_id'] = conversation_id
#
# def add_context_data(self, key: str, value):
# """添加上下文数据"""
# self.current_context[key] = value
#
# def get_context_data(self, key: str, default=None):
# """获取上下文数据"""
# return self.current_context.get(key, default)
class LLMManager:
"""Chat 管理器(与 ASRManager 结构、接口完全一致)"""
_instance: Optional[ChatBase] = None # 单例存储(对应 ASRManager._instance
@classmethod
async def initialize(cls) -> Tuple[bool, str]:
"""初始化 Chat 服务(与 ASRManager.initialize 接口一致)"""
try:
# 1. 通过工厂创建 Chat 实例
cls._instance = ChatFactory.get_chat_client()
# 2. 初始化 Chat 服务(如连接池)
init_success = await cls._instance.initialize()
if not init_success:
return False, "Chat 服务初始化失败"
# 3. 检查有效连接数
valid_conn_count = await cls._instance.get_valid_connection_count()
if valid_conn_count == 0:
return False, f"Chat 有效连接数为 0(配置池大小:{settings.CHAT_POOL_SIZE}"
return True, f"Chat 初始化成功:有效连接数 {valid_conn_count}"
except Exception as e:
return False, f"Chat 初始化失败:{str(e)}"
@classmethod
def get_instance(cls) -> Optional[ChatBase]:
"""获取全局 Chat 实例(业务代码调用,与 ASR 用法一致)"""
return cls._instance
@classmethod
async def close(cls):
"""关闭 Chat 服务(FastAPI 关闭时调用,与 ASR 一致)"""
if cls._instance:
await cls._instance.close()
cls._instance = None
print("Chat 管理器:实例和连接池已关闭")
@classmethod
def is_healthy(cls) -> bool:
"""检查 Chat 服务健康状态(与 ASR 一致)"""
return cls._instance is not None
@@ -0,0 +1,692 @@
import asyncio
import json
import websockets
import numpy as np
import sounddevice as sd
from typing import Optional, Callable, Dict, Any, List
from dataclasses import dataclass, field
import uuid
import copy
from protocols import (
EventType,
MsgType,
finish_connection,
finish_session,
receive_message,
start_connection,
start_session,
task_request,
wait_for_event,
)
# ------------------------------
# 配置常量(可根据需求调整)
# ------------------------------
DEFAULT_APPID = "7069844318"
DEFAULT_ACCESS_TOKEN = "osFMEJr20SSTWRql43cJlZkAOg7iwvxu"
DEFAULT_ENDPOINT = "wss://openspeech.bytedance.com/api/v3/tts/bidirection"
DEFAULT_VOICE_TYPE = "zh_female_gaolengyujie_emo_v2_mars_bigtts"
DEFAULT_ENCODING = "pcm"
DEFAULT_SAMPLE_RATE = 16000
@dataclass
class TTSRequest:
"""TTS请求对象(带唯一标识)"""
tts_text: str
voice_type: str = DEFAULT_VOICE_TYPE
encoding: str = DEFAULT_ENCODING
speed: float = 1.0 # 语速(字节跳动TTS支持,需服务端兼容)
stream: bool = True # 是否流式合成
request_id: str = field(default_factory=lambda: str(uuid.uuid4())) # 唯一请求ID
session_id: str = field(default_factory=lambda: str(uuid.uuid4())) # 会话ID(每个请求一个会话)
class ByteDanceTTSSocketClient:
"""字节跳动 TTS WebSocket 客户端(异步/流式/带任务队列)"""
def __init__(
self,
appid: str = DEFAULT_APPID,
access_token: str = DEFAULT_ACCESS_TOKEN,
endpoint: str = DEFAULT_ENDPOINT,
max_queue_size: int = 100
):
"""
初始化客户端
:param appid: 字节跳动APP ID
:param access_token: 访问令牌
:param endpoint: WebSocket 服务端地址
:param max_queue_size: 最大队列长度(防止内存溢出)
"""
# 基础配置
self.appid = appid
self.access_token = access_token
self.endpoint = endpoint
self.max_queue_size = max_queue_size
# WebSocket 连接状态
self.websocket: Optional[websockets.WebSocketClientProtocol] = None
self.is_connected = False
self.is_processing = False # 是否正在处理请求
self.logid: Optional[str] = None # 服务端返回的日志ID
# 异步任务队列(FIFO
self.request_queue: asyncio.Queue[TTSRequest] = asyncio.Queue(maxsize=max_queue_size)
# 回调函数定义(所有回调都带request_id,方便关联请求)
self.on_task_enqueue: Callable[[str], None] = lambda req_id: None # 任务入队回调
self.on_start: Callable[[str, Dict[str, Any]], None] = lambda req_id, data: None # 合成开始回调
self.on_audio_chunk: Callable[[str, bytes], None] = lambda req_id, chunk: None # 音频块回调(原始字节)
self.on_end: Callable[[str, Dict[str, Any]], None] = lambda req_id, data: None # 合成结束回调
self.on_error: Callable[[str, str], None] = lambda req_id, msg: None # 错误回调
self.on_queue_full: Callable[[str], None] = lambda req_id: None # 队列满回调
# 音频播放相关(支持MP3格式直接播放)
self.play_stream: Optional[sd.OutputStream] = None
self.current_req_id: Optional[str] = None
self.enable_playback: bool = True # 是否启用实时播放
# 外部回调函数(返回完整结果)
self.external_callback: Optional[Callable[[str, Dict[str, Any]], None]] = None
# 存储每个请求的完整音频数据(原始字节)
self.audio_buffers: Dict[str, List[bytes]] = {}
def _get_resource_id(self, voice_type: str) -> str:
"""根据音色类型获取资源ID(字节跳动TTS协议要求)"""
if voice_type.startswith("S_"):
return "volc.megatts.default"
return "volc.service_type.10029"
async def _create_websocket_connection(self):
"""创建WebSocket连接(内部使用)"""
headers = {
"X-Api-App-Key": self.appid,
"X-Api-Access-Key": self.access_token,
"X-Api-Resource-Id": self._get_resource_id(DEFAULT_VOICE_TYPE), # 用默认音色获取资源ID
"X-Api-Connect-Id": str(uuid.uuid4()),
}
print(f"连接到 TTS 服务端: {self.endpoint}")
self.websocket = await websockets.connect(
self.endpoint,
additional_headers=headers,
max_size=10 * 1024 * 1024 # 10MB缓冲区
)
self.is_connected = True
self.logid = self.websocket.response.headers.get("x-tt-logid")
print(f"连接成功,LogID: {self.logid}")
# 发送连接启动指令
await start_connection(self.websocket)
await wait_for_event(
self.websocket, MsgType.FullServerResponse, EventType.ConnectionStarted
)
print("TTS连接已初始化完成")
async def connect(self):
"""建立WebSocket连接(外部调用,初始化一次)"""
if not self.is_connected:
try:
await self._create_websocket_connection()
# 启动队列消费协程(后台运行)
asyncio.create_task(self._consume_queue())
except Exception as e:
error_msg = f"连接失败: {str(e)}"
print(error_msg)
raise ConnectionError(error_msg)
async def disconnect(self):
"""关闭WebSocket连接"""
if self.is_connected and self.websocket:
try:
# 发送连接结束指令
await finish_connection(self.websocket)
await wait_for_event(
self.websocket, MsgType.FullServerResponse, EventType.ConnectionFinished
)
except Exception as e:
print(f"关闭连接时异常: {str(e)}")
finally:
await self.websocket.close()
self.is_connected = False
self.websocket = None
print("已断开与TTS服务端的连接")
# 清理音频播放流
if self.play_stream:
self.play_stream.stop()
self.play_stream.close()
self.play_stream = None
def set_external_callback(self, callback: Callable[[str, Dict[str, Any]], None]):
"""设置外部回调函数,用于返回完整结果"""
self.external_callback = callback
def set_playback_enabled(self, enabled: bool):
"""设置是否启用音频实时播放"""
self.enable_playback = enabled
print(f"音频实时播放已{'启用' if enabled else '禁用'}")
async def synthesize(self, tts_text: str, **kwargs) -> str:
"""
异步非阻塞添加TTS请求到队列
:param tts_text: 要合成的文本
:param kwargs: 其他TTS参数(voice_type, encoding, speed等)
:return: 唯一请求ID
"""
# 创建请求对象(支持覆盖默认参数)
request = TTSRequest(tts_text=tts_text, **kwargs)
req_id = request.request_id
# 初始化音频缓冲区
self.audio_buffers[req_id] = []
# 异步入队(非阻塞)
try:
await self.request_queue.put(request)
self.on_task_enqueue(req_id)
print(f"请求 [{req_id[:8]}] 已加入队列,当前队列长度: {self.request_queue.qsize()}")
return req_id
except asyncio.QueueFull:
self.on_queue_full(req_id)
error_msg = f"队列已满(最大长度{self.max_queue_size}),请求 [{req_id[:8]}] 入队失败"
print(error_msg)
raise Exception(error_msg)
async def _consume_queue(self):
"""消费队列(后台协程,自动处理排队请求)"""
print("队列消费协程已启动")
while True:
try:
# 等待队列中有请求(阻塞,直到有任务)
request = await self.request_queue.get()
req_id = request.request_id
# 标记为处理中
self.is_processing = True
# print(f"\n开始处理请求 [{req_id[:8]}],剩余队列长度: {self.request_queue.qsize()}")
# 处理单个请求
await self._process_single_request(request)
# 标记任务完成(让Queue知道可以继续)
self.request_queue.task_done()
self.is_processing = False
except Exception as e:
error_msg = f"队列消费异常: {str(e)}"
print(error_msg)
self.is_processing = False
# 短暂等待,避免死循环占用CPU
await asyncio.sleep(0.1)
def _build_base_request(self, request: TTSRequest) -> Dict[str, Any]:
"""构建字节跳动TTS基础请求参数"""
aaa = {
"user": {"uid": str(uuid.uuid4())},
"namespace": "BidirectionalTTS",
"req_params": {
"speaker": request.voice_type,
"audio_params": {
"format": request.encoding,
"sample_rate": DEFAULT_SAMPLE_RATE,
"enable_timestamp": True,
},
"additions": json.dumps({"disable_markdown_filter": False}),
"speed": request.speed, # 语速参数(需服务端支持)
},
}
print('aaa', aaa)
return aaa
async def _send_text_stream(self, request: TTSRequest, session_id: str):
"""流式发送文本(逐字符发送,字节跳动TTS流式协议要求)"""
base_request = self._build_base_request(request)
text = request.tts_text.strip()
if not text:
print(f"请求 [{request.request_id[:8]}] 文本为空,跳过发送")
return
# 逐字符发送(控制发送速率,避免拥塞)
for char in text:
if not self.is_connected or not self.websocket:
raise ConnectionError("连接已断开,无法继续发送文本")
# 构建单个字符的任务请求
task_req = copy.deepcopy(base_request)
task_req["event"] = EventType.TaskRequest
task_req["req_params"]["text"] = char
# 发送任务请求
await task_request(
self.websocket,
json.dumps(task_req).encode("utf-8"),
session_id
)
# 控制发送速率(5ms/字符,可调整)
await asyncio.sleep(0.005)
# 发送会话结束指令
await finish_session(self.websocket, session_id)
print(f"请求 [{request.request_id[:8]}] 文本发送完成")
async def _handle_audio_response(self, req_id: str, session_id: str, request: TTSRequest) -> Dict[str, Any]:
"""处理服务端的流式音频响应"""
if not self.websocket:
raise ConnectionError("WebSocket连接未建立")
sample_rate = DEFAULT_SAMPLE_RATE
audio_received = False
try:
while True:
# 接收服务端消息(异步阻塞)
msg = await receive_message(self.websocket)
if msg.type == MsgType.FullServerResponse:
# 完整响应(开始/结束/错误)
if msg.event == EventType.SessionStarted:
# 会话开始回调
start_data = {
"session_id": session_id,
"sample_rate": sample_rate,
"encoding": request.encoding,
"voice_type": request.voice_type,
"logid": self.logid
}
self.on_start(req_id, start_data)
print(f"请求 [{req_id[:8]}] 合成开始")
elif msg.event == EventType.SessionFinished:
# 会话结束,退出循环
end_data = {"session_id": session_id, "message": "合成完成"}
self.on_end(req_id, end_data)
print(f"请求 [{req_id[:8]}] 合成结束")
break
elif msg.type == MsgType.AudioOnlyServer:
# 流式音频数据(原始字节)
audio_chunk = msg.payload
if audio_chunk:
audio_received = True
# 保存到缓冲区
self.audio_buffers[req_id].append(audio_chunk)
# 音频块回调
self.on_audio_chunk(req_id, audio_chunk)
# 实时播放(如果启用)
await self._play_audio_chunk(req_id, audio_chunk, sample_rate)
else:
# 未知消息类型
raise RuntimeError(f"收到未知消息类型: {msg.type}, 内容: {msg}")
# 组装完整结果
full_audio = b"".join(self.audio_buffers[req_id]) if self.audio_buffers[req_id] else b""
return {
"status": "completed",
"request_id": req_id,
"session_id": session_id,
"sample_rate": sample_rate,
"encoding": request.encoding,
"audio_data": full_audio, # 完整音频字节数据
"audio_length": len(full_audio),
"message": "合成成功" if audio_received else "合成完成但未收到音频数据"
}
except Exception as e:
error_msg = f"处理音频响应异常: {str(e)}"
self.on_error(req_id, error_msg)
return {
"status": "error",
"request_id": req_id,
"session_id": session_id,
"message": error_msg
}
async def _play_audio_chunk(self, req_id: str, chunk: bytes, sample_rate: int):
"""实时播放音频块(支持MP3格式)"""
if not self.enable_playback:
return
# 确保当前请求是正在播放的请求
if self.current_req_id is None:
self.current_req_id = req_id
if req_id != self.current_req_id:
# 切换请求时,重置播放流
if self.play_stream:
self.play_stream.stop()
self.play_stream.close()
self.current_req_id = req_id
try:
# 初始化播放流(如果未初始化)
if not self.play_stream:
self.play_stream = sd.OutputStream(
samplerate=sample_rate,
channels=1, # 单声道
dtype=np.float32
)
self.play_stream.start()
# MP3字节 → 音频数组(直接播放)
# 注意:sounddevice默认支持PCM格式,如果是MP3需要解码,这里简化处理(实际使用建议用pydub解码)
# 如需支持MP3播放,请安装 pydub: pip install pydub ffmpeg
try:
# 简化处理:假设服务端返回PCM(如果是MP3,需替换为解码逻辑)
audio_array = np.frombuffer(chunk, dtype=np.float32)
if audio_array.size > 0:
self.play_stream.write(audio_array)
except Exception as e:
print(f"音频播放异常: {str(e)},请确保音频格式正确")
except Exception as e:
print(f"播放流初始化失败: {str(e)}")
async def _process_single_request(self, request: TTSRequest):
"""处理单个TTS请求(完整流程:连接→启动会话→流式发送文本→接收音频→回调结果)"""
req_id = request.request_id
session_id = request.session_id
# 参数校验
if not request.tts_text.strip():
error_msg = "合成文本不能为空"
self.on_error(req_id, error_msg)
self._send_external_callback(req_id, {
"status": "error",
"request_id": req_id,
"session_id": session_id,
"message": error_msg
})
return
# 确保连接已建立(断开时自动重连)
if not self.is_connected:
print(f"请求 [{req_id[:8]}] 处理时连接已断开,尝试重连...")
try:
await self._create_websocket_connection()
except Exception as e:
error_msg = f"重连失败: {str(e)}"
self.on_error(req_id, error_msg)
self._send_external_callback(req_id, {
"status": "error",
"request_id": req_id,
"session_id": session_id,
"message": error_msg
})
return
try:
# 1. 启动会话
base_request = self._build_base_request(request)
start_session_req = copy.deepcopy(base_request)
start_session_req["event"] = EventType.StartSession
await start_session(
self.websocket,
json.dumps(start_session_req).encode("utf-8"),
session_id
)
# 等待会话启动成功
await wait_for_event(
self.websocket, MsgType.FullServerResponse, EventType.SessionStarted
)
# 2. 异步流式发送文本(后台任务,不阻塞接收音频)
send_task = asyncio.create_task(self._send_text_stream(request, session_id))
# 3. 接收并处理音频响应
result_data = await self._handle_audio_response(req_id, session_id, request)
# 4. 等待文本发送任务完成
await send_task
# 5. 发送外部回调
self._send_external_callback(req_id, result_data)
except Exception as e:
error_msg = f"处理请求 [{req_id[:8]}] 异常: {str(e)}"
self.on_error(req_id, error_msg)
self._send_external_callback(req_id, {
"status": "error",
"request_id": req_id,
"session_id": session_id,
"message": error_msg
})
finally:
# 清理缓冲区
if req_id in self.audio_buffers:
del self.audio_buffers[req_id]
def _send_external_callback(self, req_id: str, result_data: Dict[str, Any]):
"""发送外部回调(支持同步/异步回调函数)"""
if not self.external_callback:
return
try:
# 异步回调:直接await
if asyncio.iscoroutinefunction(self.external_callback):
asyncio.create_task(self.external_callback(req_id, result_data))
# 同步回调:在线程池中执行(避免阻塞事件循环)
else:
asyncio.get_event_loop().run_in_executor(
None, self.external_callback, req_id, result_data
)
except Exception as e:
print(f"外部回调执行异常: {str(e)}")
async def wait_all_completed(self):
"""等待队列中所有任务处理完成(阻塞)"""
await self.request_queue.join()
print("\n所有队列任务已处理完成")
# ------------------------------
# 使用示例(与你提供的风格完全一致)
# ------------------------------
class TTSManager:
"""TTS管理器 - 供外部代码调用(封装客户端,简化使用)"""
def __init__(
self,
appid: str = DEFAULT_APPID,
access_token: str = DEFAULT_ACCESS_TOKEN,
endpoint: str = DEFAULT_ENDPOINT
):
self.client = ByteDanceTTSSocketClient(
appid=appid,
access_token=access_token,
endpoint=endpoint
)
self._setup_internal_callbacks()
def _setup_internal_callbacks(self):
"""设置内部回调(日志/状态提示)"""
def on_task_enqueue(req_id: str):
"""任务入队回调"""
print(f"📥 任务 [{req_id[:8]}] 已入队")
def on_tts_start(req_id: str, data: Dict[str, Any]):
"""合成开始回调"""
print(f"🎤 合成开始 [{req_id[:8]}] - 采样率: {data['sample_rate']}, 编码: {data['encoding']}")
def on_audio_chunk(req_id: str, chunk: bytes):
"""音频块回调(内部仅打印日志,外部通过external_callback获取)"""
print(f"🔊 收到音频块 [{req_id[:8]}] - 大小: {len(chunk)}字节", end="\r")
def on_tts_end(req_id: str, data: Dict[str, Any]):
"""合成结束回调"""
print(f"\n🏁 合成结束 [{req_id[:8]}] - 会话ID: {data['session_id']}")
def on_tts_error(req_id: str, msg: str):
"""错误回调"""
print(f"\n❌ 合成失败 [{req_id[:8]}] - 错误: {msg}")
def on_queue_full(req_id: str):
"""队列满回调"""
print(f"⚠️ 队列已满,请求 [{req_id[:8]}] 入队失败")
# 绑定内部回调
self.client.on_task_enqueue = on_task_enqueue
self.client.on_start = on_tts_start
self.client.on_audio_chunk = on_audio_chunk
self.client.on_end = on_tts_end
self.client.on_error = on_tts_error
self.client.on_queue_full = on_queue_full
async def initialize(self):
"""初始化连接"""
await self.client.connect()
async def shutdown(self):
"""关闭连接"""
await self.client.disconnect()
def set_result_callback(self, callback: Callable[[str, Dict[str, Any]], None]):
"""设置外部结果回调(获取完整音频数据)"""
self.client.set_external_callback(callback)
def set_playback_enabled(self, enabled: bool):
"""设置是否启用实时播放"""
self.client.set_playback_enabled(enabled)
async def synthesize(self, text: str, **kwargs) -> str:
"""
异步非阻塞合成文本
:param text: 要合成的文本
:param kwargs: 其他参数(voice_type, encoding, speed等)
:return: 请求ID
"""
return await self.client.synthesize(text, **kwargs)
async def wait_all_completed(self):
"""等待所有任务完成"""
await self.client.wait_all_completed()
# ------------------------------
# 外部调用示例
# ------------------------------
async def external_usage_example():
"""外部代码使用示例"""
# 1. 创建TTS管理器(可替换为自己的appid和access_token
tts_manager = TTSManager(
appid=DEFAULT_APPID,
access_token=DEFAULT_ACCESS_TOKEN,
endpoint=DEFAULT_ENDPOINT
)
# 2. 设置外部结果回调(获取完整音频数据)
def handle_tts_result(req_id: str, result: Dict[str, Any]):
"""处理TTS完整结果(同步回调)"""
status = result.get("status")
if status == "completed":
audio_data = result.get("audio_data")
encoding = result.get("encoding")
audio_length = result.get("audio_length")
print(f"\n✅ 收到完整结果 [{req_id[:8]}] - 长度: {audio_length}字节, 编码: {encoding}")
# 保存音频文件
filename = f"tts_output_{req_id[:8]}.{encoding}"
with open(filename, "wb") as f:
f.write(audio_data)
print(f"💾 音频文件已保存: {filename}")
elif status == "error":
error_msg = result.get("message")
print(f"\n❌ 请求 [{req_id[:8]}] 处理失败: {error_msg}")
# 绑定外部回调
tts_manager.set_result_callback(handle_tts_result)
# 3. 设置是否启用实时播放(默认True)
tts_manager.set_playback_enabled(True)
# 4. 初始化连接
await tts_manager.initialize()
# 5. 异步提交多个TTS请求(非阻塞)
texts = [
"你好,这是字节跳动TTS的流式合成测试。",
"我支持异步非阻塞调用,多个请求可以排队处理。",
"每个请求都会返回唯一的ID,方便你跟踪结果。",
"音频数据会通过回调函数返回,支持实时播放和保存文件。",
"最后一个测试句子,演示队列的自动消费功能。"
]
req_ids = []
for i, text in enumerate(texts):
# 提交请求(非阻塞,立即返回)
req_id = await tts_manager.synthesize(
text,
voice_type=DEFAULT_VOICE_TYPE,
encoding=DEFAULT_ENCODING,
speed=1.0
)
req_ids.append(req_id)
print(f"📤 已提交请求 {i+1}: ID={req_id[:8]}")
# 模拟其他业务逻辑(无需等待TTS完成)
await asyncio.sleep(0.3)
# 6. 等待所有TTS任务完成(可选,根据业务需求决定是否等待)
await tts_manager.wait_all_completed()
# 7. 关闭连接(程序退出前调用)
await tts_manager.shutdown()
# ------------------------------
# 异步结果回调示例(高级用法)
# ------------------------------
async def async_result_callback(req_id: str, result: Dict[str, Any]):
"""异步结果回调(支持异步操作,如上传音频到服务器)"""
if result["status"] == "completed":
print(f"\n⚡ 异步处理结果 [{req_id[:8]}] - 开始上传音频...")
# 模拟异步上传操作
await asyncio.sleep(0.5)
print(f"⚡ 异步处理结果 [{req_id[:8]}] - 音频上传完成")
async def advanced_usage_example():
"""高级使用示例:异步回调 + 禁用播放 + 批量请求"""
tts_manager = TTSManager()
# 设置异步结果回调
tts_manager.set_result_callback(async_result_callback)
# 禁用实时播放(只获取音频数据)
tts_manager.set_playback_enabled(False)
await tts_manager.initialize()
# 批量提交请求(并行提交)
tasks = []
for i in range(3):
text = f"这是第{i+1}个高级测试文本,使用异步回调处理结果。"
task = tts_manager.synthesize(text, speed=0.9)
tasks.append(task)
# 并行提交所有请求
req_ids = await asyncio.gather(*tasks)
print(f"\n已并行提交 {len(req_ids)} 个请求")
# 等待所有任务完成
await tts_manager.wait_all_completed()
await tts_manager.shutdown()
if __name__ == "__main__":
try:
# 运行基础使用示例
asyncio.run(external_usage_example())
# 运行高级使用示例(取消注释)
# asyncio.run(advanced_usage_example())
except KeyboardInterrupt:
print("\n程序被用户中断")
except Exception as e:
print(f"程序异常: {str(e)}")
@@ -1,113 +1,91 @@
from fastapi import WebSocket
from typing import Dict, List, Optional
from pyexpat.errors import messages
from audio_ai_chat.config.logger import logger
# from audio_ai_chat.core.asr.factory import ASRFactory # 导入ASR工厂
# from audio_ai_chat.core.llm.factory import LLMFactory # 导入LLM工厂
# from audio_ai_chat.core.tts.factory import TTSFactory # 导入TTS工厂
from audio_ai_chat.config.settings import settings
from audio_ai_chat.utils.exceptions import ServiceCallError
from fastapi import WebSocket, WebSocketDisconnect
from audio_ai_chat.codec.ProtocolCodec import ProtocolCodec
from typing import Optional, Dict, Callable, Awaitable, List, Any, Coroutine
from typing import Dict, List, Optional, Callable,Any
from dataclasses import dataclass, field
import sys
import json
import asyncio
import websockets
import uuid
import logging
from audio_ai_chat.core.connection import ConnectionContext
from audio_ai_chat.config.logger import logger
from audio_ai_chat.core.asr.asr_manager import ASRManager
from audio_ai_chat.core.connection import ConnectionManager, ConnectionContext
from audio_ai_chat.codec.ProtocolCodec import ProtocolCodec, MessageType
from audio_ai_chat.utils.exceptions import ServiceCallError
from audio_ai_chat.core.llm.dify.dify import LLMConversation, llm_client
from audio_ai_chat.core.tts.tts_client import TTSManager
from functools import partial
# 全局WebSocket连接管理器(单例模式,确保全局统一)
class WebSocketConnectionManager:
"""WebSocket连接管理器"""
_instance: Optional["WebSocketConnectionManager"] = None
def __init__(self):
# 活跃连接列表
self.active_connections: List[WebSocket] = []
# 关键映射:client_id -> ConnectionContext(快速获取用户专属上下文)
self.client_context_map: Dict[str, ConnectionContext] = {}
# self.asr_client = ASRFactory.get_asr_client()
# self.tts_client = TTSFactory.get_tts_client()
# 用户LLM会话存储
# self.user_llm_conversations: Dict[str, LLMConversation] = {}
# 全局唤醒事件
self.consume_wakeup = asyncio.Event()
self.connection_manager = None
async def connect(self, client_id, websocket: WebSocket) -> ConnectionContext:
"""
建立连接+身份校验(前端主动发送身份信息)
超时逻辑:5秒内未收到前端身份信息,自动关闭连接
返回:校验通过的 ConnectionContext(保证非空)
"""
# 1. 接受连接并加入活跃列表
def __new__(cls):
if cls._instance is None:
cls._instance = super().__new__(cls)
cls._instance.active_connections: List[WebSocket] = []
cls._instance.client_context_map: Dict[str, ConnectionContext] = {}
cls._instance.connection_manager: Optional[ConnectionManager] = None
cls._instance.consume_wakeup = asyncio.Event()
return cls._instance
async def initialize(self):
"""初始化:获取ConnectionManager全局单例"""
self.connection_manager = await ConnectionManager.get_instance()
logger.info("WebSocketConnectionManager 初始化成功")
async def connect(self, client_id: str, websocket: WebSocket) -> ConnectionContext:
"""建立连接+身份校验"""
await websocket.accept()
context = ConnectionContext(client_id=client_id) # 提前创建上下文(保证最终返回非空)
self.active_connections.append(websocket)
logger.info(
f"连接 {client_id} 已接受,等待前端发送身份信息(5秒超时)...,当前连接数: {len(self.active_connections)}")
f"连接 {client_id} 已接受,等待身份信息(5秒超时),当前连接数: {len(self.active_connections)}"
)
# 2. 超时控制:5秒内未收到身份信息 -> 关闭连接
# 超时接收身份包
try:
ping_packet = await asyncio.wait_for(
websocket.receive_bytes(),
timeout=5.0
)
ping_packet = await asyncio.wait_for(websocket.receive_bytes(), timeout=5.0)
except asyncio.TimeoutError:
error_msg = f"连接 {client_id} 身份校验超时5秒未收到消息)"
error_msg = f"连接 {client_id} 身份校验超时"
logger.warning(error_msg)
# 发送超时错误响应(二进制格式)
error_packet = ProtocolCodec.pack(
MessageType.ERROR,
{"code": 1008, "message": "身份校验超时,请重试"}
MessageType.ERROR, {"code": 1008, "message": "身份校验超时,请重试"}
)
await websocket.send_bytes(error_packet)
raise TimeoutError(error_msg) # 抛出异常,进入后续清理逻辑
raise TimeoutError(error_msg)
# 3. 解包并验证包类型
print('ping_packet', ping_packet)
# 解包并验证包类型
msg_type, _, identity_data = ProtocolCodec.unpack(ping_packet)
if msg_type != MessageType.IDENTITY:
error_msg = f"连接 {client_id} 首个包类型错误(期望{MessageType.IDENTITY.value},实际{msg_type.value}"
error_msg = f"连接 {client_id} 首个包类型错误"
logger.error(error_msg)
# 发送类型错误响应
error_packet = ProtocolCodec.pack(
MessageType.ERROR,
{"code": 4001, "message": "非法请求:首个包必须是身份校验包"}
MessageType.ERROR, {"code": 4001, "message": "非法请求:首个包必须是身份校验包"}
)
await websocket.send_bytes(error_packet)
raise ValueError(error_msg)
# 4. 身份信息并校验
# todo
# 提取核心字段(必选字段校验)
# 校验身份信息
user_id = identity_data.get("user_id")
token = identity_data.get("token")
name = identity_data.get("name") or f"用户{user_id}" # 提供默认名称
name = identity_data.get("name") or f"用户{user_id}"
if not all([user_id, token]):
error_msg = f"连接 {client_id} 身份信息不完整(缺少user_id或token"
error_msg = f"连接 {client_id} 身份信息不完整"
logger.error(error_msg)
error_packet = ProtocolCodec.pack(
MessageType.ERROR,
{"code": 4003, "message": "身份信息不完整:必须包含user_id和token"}
MessageType.ERROR, {"code": 4003, "message": "身份信息不完整:必须包含user_id和token"}
)
await websocket.send_bytes(error_packet)
raise ValueError(error_msg)
# TODO: 实际身份校验逻辑(根据你的业务扩展)
# 5. 校验通过:更新上下文并响应前端
# 创建/获取连接上下文
context = await self.connection_manager.create_or_reconnect_context(
new_client_id=client_id, user_id=user_id
)
context.set_user_info(token, user_id, name)
self.client_context_map[client_id] = context # 加入上下文映射
self.client_context_map[client_id] = context
# 发送成功响应
# 响应身份校验成功
success_packet = ProtocolCodec.pack(
MessageType.IDENTITY,
{
@@ -117,309 +95,225 @@ class WebSocketConnectionManager:
}
)
await websocket.send_bytes(success_packet)
logger.info(f"用户 {user_id}{name})身份校验通过,连接就绪client_id: {client_id}")
logger.info(f"用户 {user_id}{name})身份校验通过(client_id: {client_id}")
return context
# 初始化LLM会话
# self._init_llm_conversation(user_id)
# return conn_id, user_id, conn_id # conn_id 同时作为 tts_session_id
# def _init_llm_conversation(self, user_id: str):
# """初始化用户LLM会话"""
# if user_id not in self.user_llm_conversations:
# self.user_llm_conversations[user_id] = LLMConversation(
# user_id=user_id,
# scene_description="语音识别对话场景"
# )
def disconnect(self, websocket: WebSocket, conn_id: str):
"""断开连接并清理资源"""
if websocket in self.active_connections:
self.active_connections.remove(websocket)
logger.info(f"连接 {conn_id} 已断开,当前连接数: {len(self.active_connections)}")
# async def setup_tts_manager(self, result_queue: asyncio.Queue) -> TTSManager:
# """初始化TTS管理器"""
#
# def handle_tts_result(req_id: str, result: Dict[str, Any]):
# """TTS结果回调处理"""
# try:
# status = result.get("status")
# if status == "completed":
# audio_data = result.get("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(pcm_bytes)
# except Exception as e:
# logger.error(f"TTS结果处理失败: {str(e)}")
#
# self.tts_client = TTSFactory.get_tts_client()
# tts_manager = TTSManager()
# tts_manager.set_result_callback(handle_tts_result)
# tts_manager.set_playback_enabled(False)
# await tts_manager.initialize()
# return tts_manager
async def asr_result_callback(self, result: dict, websocket: WebSocket,
user_id: str, result_queue: asyncio.Queue):
"""ASR结果回调处理"""
try:
logger.info(f"ASR识别结果: {result}")
final_asr_text = result.get("text", "").strip()
# 转发ASR结果到前端队列
if final_asr_text:
print(f"插入ASR结果时队列大小: {result_queue.qsize()}")
self.consume_wakeup.set() # 唤醒消费协程
# 异步调用大模型
llm_conversation = self.user_llm_conversations.get(user_id)
if llm_conversation:
asyncio.create_task(
self.call_llm_and_send(
query=final_asr_text,
conversation=llm_conversation,
websocket=websocket
)
def _create_asr_callback(self, context: ConnectionContext) -> Callable[[dict], None]:
"""
闭包:为当前连接创建专属的ASR回调函数
回调内部持有ConnectionContext引用,直接操作其消息队列
"""
async def asr_result_callback(result: dict):
"""专属回调:将ASR结果打包后插入当前连接的消息队列"""
try:
print('result', 'result', result)
# 处理错误结果
if result.get("error"):
logger.error(f"ASR错误(client_id: {context.client_id}):{result['error']}")
# 打包错误消息
error_packet = ProtocolCodec.pack(
MessageType.ERROR,
{"code": 5001, "message": f"ASR服务错误:{result['error']}"}
)
context.message_queue.put_nowait(error_packet)
else:
# 3. 调用Dify流式接口(传入ASR元数据,用于存储到对话历史)
print('context', context.user_id)
final_asr_text = result.get("text", "")
# 2. 调用大模型(异步)
if final_asr_text:
asyncio.create_task(
self.call_llm_and_send(
context=context,
query=final_asr_text,
conversation=context.llm_session
)
)
logger.debug(
f"ASR结果入队(client_id: {context.client_id}):"
f"文本={result['text']},最终结果={result['is_final']}"
)
except Exception as e:
logger.error(f"ASR回调处理失败(client_id: {context.client_id}):{str(e)}")
return asr_result_callback
# ====================== 调用大模型 ======================
# ====================== 大模型流式回调 ======================
@staticmethod
async def llm_stream_callback(context, chunk: str, conversation_id: str, is_finished: bool):
"""大模型流式回调(纯异步,无阻塞)"""
if not chunk:
return
print('大模型流式回调', chunk)
req_id = await context.tts_client.synthesize(chunk)
print('req_id', req_id)
async def call_llm_and_send(self,context ,query: str, conversation: LLMConversation):
"""调用大模型,流式结果转发前端 + TTS"""
if not query:
return
logger.info(f"调用大模型 - 用户(): {query}")
try:
stream_callback = partial(WebSocketConnectionManager.llm_stream_callback, context)
conv_id, full_reply = await llm_client.send_message(
query=query,
conversation=conversation,
stream_callback=stream_callback, # 传递绑定后的回调
response_mode="streaming"
)
logger.info(f"大模型回复完成 - 会话ID: {conv_id}, 完整回复: {full_reply}")
except Exception as e:
logger.error(f"ASR回调执行失败: {str(e)}")
logger.error(f"大模型调用失败: {str(e)}")
# async def llm_stream_callback(self, chunk: str, tts_manager: TTSManager):
# """大模型流式回调处理"""
# if not chunk:
# return
# try:
# # 提交TTS合成请求
# await tts_manager.synthesize(chunk)
# await asyncio.sleep(0) # 让出调度权
# except Exception as e:
# logger.error(f"LLM流式回调处理失败: {str(e)}")
# async def call_llm_and_send(self, query: str, conversation: LLMConversation, websocket: WebSocket):
# """调用大模型并处理结果"""
# logger.info(f"调用大模型 - 用户({conversation.user_id}): {query}")
# try:
# conv_id, full_reply = await llm_client.send_message(
# query=query,
# conversation=conversation,
# stream_callback=self.llm_stream_callback,
# response_mode="streaming"
# )
# logger.info(f"大模型回复完成 - 会话ID: {conv_id}, 完整回复: {full_reply}")
# except Exception as e:
# logger.error(f"大模型调用失败: {str(e)}")
# if not websocket.client_state.disconnected:
# await websocket.send_json({
# "type": "llm_error",
# "data": {"error": str(e)}
# })
async def recv_frontend_data(self, websocket: WebSocket, asr_conn):
"""接收前端音频数据并推送到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:
# logger.warning("音频数据插入ASR失败(队列满/连接失效)")
except WebSocketDisconnect:
logger.info("前端主动断开连接")
asr_conn.stop_event.set()
break
except Exception as e:
logger.error(f"接收前端数据失败: {str(e)}")
# asr_conn.stop_event.set()
# if not websocket.client_state.disconnected:
# await websocket.send_json({"error": f"接收数据失败: {str(e)}"})
# break
async def send_results(self, websocket: WebSocket, result_queue: asyncio.Queue, asr_conn):
"""从结果队列发送数据到前端"""
while True:
try:
# 等待队列数据或超时
result = await asyncio.wait_for(result_queue.get(), timeout=0.05)
# if not websocket.client_state.disconnected:
await websocket.send_bytes(result)
except asyncio.TimeoutError:
if asr_conn.stop_event.is_set():
break
continue
except Exception as e:
logger.error(f"发送结果到前端失败: {str(e)}")
asr_conn.stop_event.set()
break
async def handle_connection(self, websocket: WebSocket):
"""处理单个WebSocket连接的完整生命周期"""
client_id = str(id(websocket))
context = None
client_id = str(id(websocket)) # 生成唯一连接ID
logger.info(f"新WebSocket连接:client_id={client_id}")
context: Optional[ConnectionContext] = None
asr_conn = None
llm_conn = None # 新增:LLM连接变量
try:
# 1. 建立连接并获取上下文
context = await self.connect(client_id, websocket)
if not context:
logger.error(f"连接 {client_id} 上下文创建失败")
return
# 接收前端数据
# 2. 获取ASR连接和专属回调
asr_client = ASRManager.get_instance()
asr_conn = await asr_client.get_connection()
if not asr_conn:
raise ServiceCallError("获取ASR连接失败")
# 创建当前连接的专属ASR回调(闭包绑定context)
asr_callback = self._create_asr_callback(context)
# 3. 启动ASR通信任务(传入专属回调)
communication_task = asyncio.create_task(
asr_client.start_communication(asr_conn, asr_callback)
)
context.llm_session = LLMConversation(
user_id=context.user_id,
scene_description="语音识别对话场景"
)
context.tts_client = TTSManager()
def handle_tts_result(context, req_id: str, result: Dict[str, Any]):
"""处理TTS结果回调"""
status = result.get("status")
print('处理TTS结果回调')
if status == "completed":
audio_data = result.get("audio_data")
pack_data = ProtocolCodec.pack(MessageType.AUDIO_DATA, audio_data)
context.message_queue.put_nowait(pack_data)
print('插入', len(audio_data))
stream_callback = partial(handle_tts_result, context)
context.tts_client.set_result_callback(stream_callback)
# 3. 设置是否播放(可选,默认True)
context.tts_client.set_playback_enabled(False) # 设置为False则不播放
# 4. 初始化连接
await context.tts_client.initialize()
# 4. 定义前端数据接收任务
async def recv_frontend_data():
"""接收前端音频/控制指令"""
# while not asr_conn.stop_event.is_set():
while True:
try:
if not context.message_queue.empty():
await asyncio.sleep(0) # 立即让权
continue
raw_bytes = await websocket.receive_bytes()
unpack_bytes = ProtocolCodec.unpack(raw_bytes)
success = await push_audio_data(asr_conn, unpack_bytes)
# if not success:
# print("音频数据插入失败(队列满/连接失效)")
msg_type, sequence, data = ProtocolCodec.unpack(raw_bytes)
if msg_type == MessageType.AUDIO_DATA:
# 推送音频数据到ASR
success = await asr_client.push_audio(asr_conn, data)
if not success:
logger.warning(f"连接 {client_id} 音频推送失败(队列满/连接失效)")
elif msg_type == MessageType.CONTROL:
# 处理控制指令(如暂停/继续ASR
logger.info(f"连接 {client_id} 收到控制指令:{data}")
if data.get("action") == "stop_asr":
asr_conn.stop_event.set()
else:
logger.warning(f"连接 {client_id} 收到未知消息类型:{msg_type.value}")
except WebSocketDisconnect:
logger.info(f"前端 {conn_id} 主动断开连接")
asr_conn.stop_event.set()
logger.info(f"前端 {client_id} 主动断开连接")
break
except Exception as e:
logger.error(f"接收前端数据失败: {str(e)}")
asr_conn.stop_event.set()
await websocket.send_json({"error": f"接收数据失败: {str(e)}"})
logger.error(f"连接 {client_id} 接收前端数据失败{str(e)}")
break
# 发送 ASR 结果
# 5. 定义ASR结果发送任务(从上下文队列取数据)
async def send_asr_result():
"""从结果队列发送 ASR 结果到前端(二进制格式)"""
while True:
try:
result = await asyncio.wait_for(context.message_queue.get(), timeout=0.05)
await websocket.send_bytes(result)
# 从当前连接的消息队列获取ASR结果(超时0.05秒避免阻塞)
result_packet = await asyncio.wait_for(
context.message_queue.get(), timeout=0.05
)
print('发送', result_packet)
await websocket.send_bytes(result_packet)
except asyncio.TimeoutError:
continue
continue # 无数据时继续等待
except Exception as e:
logger.error(f"发送 ASR 结果失败: {str(e)}")
logger.error(f"连接 {client_id} 发送ASR结果失败{str(e)}")
break
# 6. 启动任务并等待完成
task_send = asyncio.create_task(send_asr_result())
task_recv = asyncio.create_task(recv_frontend_data())
try:
# 等待两个任务,只要有一个完成就返回(比如前端断开/发送出错)
done, pending = await asyncio.wait(
[task_recv, task_send],
return_when=asyncio.FIRST_COMPLETED,
timeout=None # 无限等待,直到有任务完成
)
finally:
# 确保协程正确退出
# 等待剩余任务完成
for task in pending:
task.cancel()
await asyncio.gather(task_recv, task_send, return_exceptions=True)
pass
done, pending = await asyncio.wait(
[task_recv, task_send, communication_task],
return_when=asyncio.FIRST_COMPLETED
)
# 取消未完成的任务
for task in pending:
task.cancel()
await asyncio.gather(*pending, return_exceptions=True)
except Exception as e:
logger.error(f"WebSocket连接处理异常: {str(e)}")
if websocket in self.active_connections:
self.active_connections.remove(websocket)
# 移除上下文映射(如果已添加)
if context is not None and context.client_id in self.client_context_map:
del self.client_context_map[client_id]
logger.error(f"连接 {client_id} 处理异常{str(e)}")
# 异常时发送错误消息给前端
if websocket.state == "CONNECTED":
error_packet = ProtocolCodec.pack(
MessageType.ERROR, {"code": 5000, "message": f"服务异常:{str(e)}"}
)
await websocket.send_bytes(error_packet)
finally:
# 6. 统一资源清理(无论成功/失败,都执行
# 关闭WebSocket连接
try:
if hasattr(websocket, "state") and websocket.state == "CONNECTED":
await websocket.close(code=1008, reason="连接终止")
except Exception as close_e:
logger.warning(f"关闭连接失败 (client_id: {client_id}): {str(close_e)}")
# 7. 资源清理(关键
# 停止ASR通信任务
# if communication_task and not communication_task.done():
# communication_task.cancel()
# try:
# await communication_task
# except Exception as e:
# logger.warning(f"连接 {client_id} ASR任务取消异常:{str(e)}")
# 移除活跃连接
# 关闭ASR连接
# if asr_conn:
# await asr_client.close_connection(asr_conn)
# 关闭WebSocket连接
if websocket.state == "CONNECTED":
await websocket.close(code=1008, reason="连接终止")
# 移除连接和上下文
if websocket in self.active_connections:
self.active_connections.remove(websocket)
# 移除上下文映射
if client_id in self.client_context_map:
del self.client_context_map[client_id]
if context:
pass
logger.info(f"连接资源清理完成 (client_id: {client_id}),当前连接数: {len(self.active_connections)}")
# 1. 建立连接
# 2. 初始化TTS
# tts_manager = await self.setup_tts_manager(result_queue)
# 3. 获取ASR连接
asr_conn = await get_idle_asr_connection()
if not asr_conn:
await websocket.send_json({"error": "ASR服务暂时不可用", "text": ""})
return
# 4. 启动ASR通信协程
# asr_callback = lambda res: self.asr_result_callback(res, websocket, user_id, result_queue)
# asr_task = asyncio.create_task(handle_asr_communication(asr_conn, asr_callback))
#
# # 5. 启动数据接收和发送协程
# task_recv = asyncio.create_task(self.recv_frontend_data(websocket, asr_conn))
# task_send = asyncio.create_task(self.send_results(websocket, result_queue, asr_conn))
#
# # 6. 等待任一任务完成
# done, pending = await asyncio.wait(
# [task_recv, task_send],
# return_when=asyncio.FIRST_COMPLETED
# )
# except Exception as e:
# logger.error(f"WebSocket连接处理异常: {str(e)}")
# if asr_conn:
# asr_conn.stop_event.set()
# if not websocket.client_state.disconnected:
# await websocket.send_json({"error": str(e)})
# logger.error(f"连接 {client_id} 建立失败: {type(e).__name__}: {e}")
# try:
# # 确保连接已关闭(处理未正常关闭的情况)
# if websocket.client_state == "CONNECTED": # 根据实际WebSocket类型调整状态判断
# await websocket.close(code=1008, reason=str(e))
# except:
# pass
# 移除活跃连接(避免内存泄漏)
# if websocket in self.active_connections:
# self.active_connections.remove(websocket)
# # 移除上下文映射(如果已添加)
# if context is not None and context.client_id in self.client_context_map:
# del self.client_context_map[client_id]
# finally:
# pass
# 7. 资源清理
# logger.info(f"开始清理连接 {conn_id} 的资源")
# # 停止ASR
# 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
#
# # 清理TTS
# if tts_manager:
# await tts_manager.cleanup() # 假设TTSManager有cleanup方法,无则忽略
#
# # 断开连接
# if websocket:
# self.disconnect(websocket, conn_id)
# try:
# await websocket.close()
# except Exception:
# pass
#
# logger.info(f"连接 {conn_id} 资源清理完成")
logger.info(
f"连接 {client_id} 资源清理完成,当前连接数: {len(self.active_connections)}"
)
@@ -0,0 +1,449 @@
from fastapi import WebSocket
from typing import Dict, List, Optional
from pyexpat.errors import messages
from audio_ai_chat.config.logger import logger
from audio_ai_chat.core.asr.factory import ASRFactory # 导入ASR工厂
# from audio_ai_chat.core.llm.factory import LLMFactory # 导入LLM工厂
# from audio_ai_chat.core.tts.factory import TTSFactory # 导入TTS工厂
from audio_ai_chat.config.settings import settings
from audio_ai_chat.utils.exceptions import ServiceCallError
from fastapi import WebSocket, WebSocketDisconnect
from audio_ai_chat.codec.ProtocolCodec import ProtocolCodec
from typing import Optional, Dict, Callable, Awaitable, List, Any, Coroutine
from dataclasses import dataclass, field
import sys
import json
import asyncio
import websockets
import uuid
import logging
from audio_ai_chat.core.connection import ConnectionManager,ConnectionContext
from audio_ai_chat.codec.ProtocolCodec import ProtocolCodec, MessageType
from audio_ai_chat.core.asr.asr_manager import ASRManager
async def asr_result_callback(result: dict):
if result.get("error"):
print(f"ASR错误:{result['error']}")
else:
print(f"ASR结果:{result['text']}(最终结果:{result['is_final']}")
class WebSocketConnectionManager:
"""WebSocket连接管理器"""
def __init__(self):
# 活跃连接列表
self.active_connections: List[WebSocket] = []
# 关键映射:client_id -> ConnectionContext(快速获取用户专属上下文)
self.client_context_map: Dict[str, ConnectionContext] = {}
self.connection_manager: Optional[ConnectionManager] = None
# self.asr_client = ASRFactory.get_asr_client()
# self.tts_client = TTSFactory.get_tts_client()
# 用户LLM会话存储
# self.user_llm_conversations: Dict[str, LLMConversation] = {}
# 全局唤醒事件
self.consume_wakeup = asyncio.Event()
async def initialize(self):
"""初始化:获取ConnectionManager全局单例(在FastAPI启动时调用)"""
self.connection_manager = await ConnectionManager.get_instance()
logger.info("WebSocketConnectionManager 初始化成功(绑定全局ConnectionManager")
async def connect(self, client_id, websocket: WebSocket) -> ConnectionContext:
"""
建立连接+身份校验(前端主动发送身份信息)
超时逻辑:5秒内未收到前端身份信息,自动关闭连接
返回:校验通过的 ConnectionContext(保证非空)
"""
# 1. 接受连接并加入活跃列表
await websocket.accept()
self.active_connections.append(websocket)
logger.info(
f"连接 {client_id} 已接受,等待前端发送身份信息(5秒超时)...,当前连接数: {len(self.active_connections)}")
# 2. 超时控制:5秒内未收到身份信息 -> 关闭连接
try:
ping_packet = await asyncio.wait_for(
websocket.receive_bytes(),
timeout=5.0
)
except asyncio.TimeoutError:
error_msg = f"连接 {client_id} 身份校验超时(5秒未收到消息)"
logger.warning(error_msg)
# 发送超时错误响应(二进制格式)
error_packet = ProtocolCodec.pack(
MessageType.ERROR,
{"code": 1008, "message": "身份校验超时,请重试"}
)
await websocket.send_bytes(error_packet)
raise TimeoutError(error_msg) # 抛出异常,进入后续清理逻辑
# 3. 解包并验证包类型
print('ping_packet', ping_packet)
msg_type, _, identity_data = ProtocolCodec.unpack(ping_packet)
if msg_type != MessageType.IDENTITY:
error_msg = f"连接 {client_id} 首个包类型错误(期望{MessageType.IDENTITY.value},实际{msg_type.value}"
logger.error(error_msg)
# 发送类型错误响应
error_packet = ProtocolCodec.pack(
MessageType.ERROR,
{"code": 4001, "message": "非法请求:首个包必须是身份校验包"}
)
await websocket.send_bytes(error_packet)
raise ValueError(error_msg)
# 4. 身份信息并校验
# todo
# 提取核心字段(必选字段校验)
user_id = identity_data.get("user_id")
token = identity_data.get("token")
name = identity_data.get("name") or f"用户{user_id}" # 提供默认名称
if not all([user_id, token]):
error_msg = f"连接 {client_id} 身份信息不完整(缺少user_id或token"
logger.error(error_msg)
error_packet = ProtocolCodec.pack(
MessageType.ERROR,
{"code": 4003, "message": "身份信息不完整:必须包含user_id和token"}
)
await websocket.send_bytes(error_packet)
raise ValueError(error_msg)
# TODO: 实际身份校验逻辑(根据你的业务扩展)
# 5. 校验通过:更新上下文并响应前端
context = await self.connection_manager.create_or_reconnect_context(
new_client_id=client_id,
user_id=user_id
)
context = ConnectionContext(client_id=client_id) # 提前创建上下文(保证最终返回非空)
context.set_user_info(token, user_id, name)
self.client_context_map[client_id] = context # 加入上下文映射
# 发送成功响应
success_packet = ProtocolCodec.pack(
MessageType.IDENTITY,
{
"code": 200,
"message": "身份校验成功,连接已就绪",
"data": {"client_id": client_id, "user_id": user_id, "name": name}
}
)
await websocket.send_bytes(success_packet)
logger.info(f"用户 {user_id}{name})身份校验通过,连接就绪(client_id: {client_id}")
return context
# 初始化LLM会话
# self._init_llm_conversation(user_id)
# return conn_id, user_id, conn_id # conn_id 同时作为 tts_session_id
# def _init_llm_conversation(self, user_id: str):
# """初始化用户LLM会话"""
# if user_id not in self.user_llm_conversations:
# self.user_llm_conversations[user_id] = LLMConversation(
# user_id=user_id,
# scene_description="语音识别对话场景"
# )
def disconnect(self, websocket: WebSocket, conn_id: str):
"""断开连接并清理资源"""
if websocket in self.active_connections:
self.active_connections.remove(websocket)
logger.info(f"连接 {conn_id} 已断开,当前连接数: {len(self.active_connections)}")
# async def setup_tts_manager(self, result_queue: asyncio.Queue) -> TTSManager:
# """初始化TTS管理器"""
#
# def handle_tts_result(req_id: str, result: Dict[str, Any]):
# """TTS结果回调处理"""
# try:
# status = result.get("status")
# if status == "completed":
# audio_data = result.get("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(pcm_bytes)
# except Exception as e:
# logger.error(f"TTS结果处理失败: {str(e)}")
#
# self.tts_client = TTSFactory.get_tts_client()
# tts_manager = TTSManager()
# tts_manager.set_result_callback(handle_tts_result)
# tts_manager.set_playback_enabled(False)
# await tts_manager.initialize()
# return tts_manager
# async def asr_result_callback(self, result: dict, websocket: WebSocket,
# user_id: str, result_queue: asyncio.Queue):
# """ASR结果回调处理"""
# try:
# logger.info(f"ASR识别结果: {result}")
# final_asr_text = result.get("text", "").strip()
#
# # 转发ASR结果到前端队列
# if final_asr_text:
# print(f"插入ASR结果时队列大小: {result_queue.qsize()}")
# self.consume_wakeup.set() # 唤醒消费协程
#
# # 异步调用大模型
# llm_conversation = self.user_llm_conversations.get(user_id)
# if llm_conversation:
# asyncio.create_task(
# self.call_llm_and_send(
# query=final_asr_text,
# conversation=llm_conversation,
# websocket=websocket
# )
# )
# except Exception as e:
# logger.error(f"ASR回调执行失败: {str(e)}")
# async def llm_stream_callback(self, chunk: str, tts_manager: TTSManager):
# """大模型流式回调处理"""
# if not chunk:
# return
# try:
# # 提交TTS合成请求
# await tts_manager.synthesize(chunk)
# await asyncio.sleep(0) # 让出调度权
# except Exception as e:
# logger.error(f"LLM流式回调处理失败: {str(e)}")
# async def call_llm_and_send(self, query: str, conversation: LLMConversation, websocket: WebSocket):
# """调用大模型并处理结果"""
# logger.info(f"调用大模型 - 用户({conversation.user_id}): {query}")
# try:
# conv_id, full_reply = await llm_client.send_message(
# query=query,
# conversation=conversation,
# stream_callback=self.llm_stream_callback,
# response_mode="streaming"
# )
# logger.info(f"大模型回复完成 - 会话ID: {conv_id}, 完整回复: {full_reply}")
# except Exception as e:
# logger.error(f"大模型调用失败: {str(e)}")
# if not websocket.client_state.disconnected:
# await websocket.send_json({
# "type": "llm_error",
# "data": {"error": str(e)}
# })
async def recv_frontend_data(self, websocket: WebSocket, asr_conn):
"""接收前端音频数据并推送到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:
# logger.warning("音频数据插入ASR失败(队列满/连接失效)")
except WebSocketDisconnect:
logger.info("前端主动断开连接")
asr_conn.stop_event.set()
break
except Exception as e:
logger.error(f"接收前端数据失败: {str(e)}")
# asr_conn.stop_event.set()
# if not websocket.client_state.disconnected:
# await websocket.send_json({"error": f"接收数据失败: {str(e)}"})
# break
async def send_results(self, websocket: WebSocket, result_queue: asyncio.Queue, asr_conn):
"""从结果队列发送数据到前端"""
while True:
try:
# 等待队列数据或超时
result = await asyncio.wait_for(result_queue.get(), timeout=0.05)
# if not websocket.client_state.disconnected:
await websocket.send_bytes(result)
except asyncio.TimeoutError:
if asr_conn.stop_event.is_set():
break
continue
except Exception as e:
logger.error(f"发送结果到前端失败: {str(e)}")
asr_conn.stop_event.set()
break
async def handle_connection(self, websocket: WebSocket):
"""处理单个WebSocket连接的完整生命周期"""
new_client_id = str(id(websocket))
logger.info(f"新WebSocket连接:client_id={new_client_id}")
context = None
try:
context = await self.connect(new_client_id, websocket)
asr_client = ASRManager.get_instance()
asr_conn = await asr_client.get_connection()
if not asr_conn:
print("获取ASR连接失败")
raise
communication_task = asyncio.create_task(
asr_client.start_communication(asr_conn, asr_result_callback)
)
# 接收前端数据
async def recv_frontend_data():
"""接收前端音频/控制指令"""
# while not asr_conn.stop_event.is_set():
while True:
try:
if not context.message_queue.empty():
await asyncio.sleep(0) # 立即让权
continue
raw_bytes = await websocket.receive_bytes()
msg_type, sequence, unpack_bytes = ProtocolCodec.unpack(raw_bytes)
if msg_type == MessageType.AUDIO_DATA:
success = await asr_client.push_audio(asr_conn, unpack_bytes)
if not success:
print(f"音频数据插入失败(队列满/连接失效)")
else:
print(f"其他类型数据", msg_type)
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 True:
try:
result = await asyncio.wait_for(context.message_queue.get(), timeout=0.05)
await websocket.send_bytes(result)
except asyncio.TimeoutError:
continue
except Exception as e:
logger.error(f"发送 ASR 结果失败: {str(e)}")
break
task_send = asyncio.create_task(send_asr_result())
task_recv = asyncio.create_task(recv_frontend_data())
try:
# 等待两个任务,只要有一个完成就返回(比如前端断开/发送出错)
done, pending = await asyncio.wait(
[task_recv, task_send],
return_when=asyncio.FIRST_COMPLETED,
timeout=None # 无限等待,直到有任务完成
)
finally:
# 确保协程正确退出
# 等待剩余任务完成
for task in pending:
task.cancel()
await asyncio.gather(task_recv, task_send, return_exceptions=True)
pass
except Exception as e:
logger.error(f"WebSocket连接处理异常: {str(e)}")
if websocket in self.active_connections:
self.active_connections.remove(websocket)
# 移除上下文映射(如果已添加)
if context is not None and context.client_id in self.client_context_map:
del self.client_context_map[new_client_id]
finally:
# 6. 统一资源清理(无论成功/失败,都执行)
# 关闭WebSocket连接
try:
if hasattr(websocket, "state") and websocket.state == "CONNECTED":
await websocket.close(code=1008, reason="连接终止")
except Exception as close_e:
logger.warning(f"关闭连接失败 (client_id: {new_client_id}): {str(close_e)}")
# 移除活跃连接
if websocket in self.active_connections:
self.active_connections.remove(websocket)
# 移除上下文映射
if new_client_id in self.client_context_map:
del self.client_context_map[new_client_id]
if context:
pass
logger.info(f"连接资源清理完成 (client_id: {new_client_id}),当前连接数: {len(self.active_connections)}")
# 1. 建立连接
# 2. 初始化TTS
# tts_manager = await self.setup_tts_manager(result_queue)
# 3. 获取ASR连接
# asr_conn = await get_idle_asr_connection()
# if not asr_conn:
# await websocket.send_json({"error": "ASR服务暂时不可用", "text": ""})
# return
# 4. 启动ASR通信协程
# asr_callback = lambda res: self.asr_result_callback(res, websocket, user_id, result_queue)
# asr_task = asyncio.create_task(handle_asr_communication(asr_conn, asr_callback))
#
# # 5. 启动数据接收和发送协程
# task_recv = asyncio.create_task(self.recv_frontend_data(websocket, asr_conn))
# task_send = asyncio.create_task(self.send_results(websocket, result_queue, asr_conn))
#
# # 6. 等待任一任务完成
# done, pending = await asyncio.wait(
# [task_recv, task_send],
# return_when=asyncio.FIRST_COMPLETED
# )
# except Exception as e:
# logger.error(f"WebSocket连接处理异常: {str(e)}")
# if asr_conn:
# asr_conn.stop_event.set()
# if not websocket.client_state.disconnected:
# await websocket.send_json({"error": str(e)})
# logger.error(f"连接 {client_id} 建立失败: {type(e).__name__}: {e}")
# try:
# # 确保连接已关闭(处理未正常关闭的情况)
# if websocket.client_state == "CONNECTED": # 根据实际WebSocket类型调整状态判断
# await websocket.close(code=1008, reason=str(e))
# except:
# pass
# 移除活跃连接(避免内存泄漏)
# if websocket in self.active_connections:
# self.active_connections.remove(websocket)
# # 移除上下文映射(如果已添加)
# if context is not None and context.client_id in self.client_context_map:
# del self.client_context_map[client_id]
# finally:
# pass
# 7. 资源清理
# logger.info(f"开始清理连接 {conn_id} 的资源")
# # 停止ASR
# 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
#
# # 清理TTS
# if tts_manager:
# await tts_manager.cleanup() # 假设TTSManager有cleanup方法,无则忽略
#
# # 断开连接
# if websocket:
# self.disconnect(websocket, conn_id)
# try:
# await websocket.close()
# except Exception:
# pass
#
# logger.info(f"连接 {conn_id} 资源清理完成")
+19 -2
View File
@@ -8,18 +8,35 @@ from audio_ai_chat.codec.ProtocolCodec import ProtocolCodec # 已有加密类
from audio_ai_chat.utils.exceptions import CodecError, ServiceCallError
from contextlib import asynccontextmanager
from frontend_ws import frontend_websocket_handler
from audio_ai_chat.core.asr.asr_manager import ASRManager
# FastAPI 启动时初始化 ASR 连接池
@asynccontextmanager
async def lifespan(app: FastAPI):
# 启动时执行(原 startup 逻辑)
print(' FastAPI 启动时初始化 ASR 连接池')
# await init_asr_pool()
yield # 应用运行中
# 关闭时执行(可选,比如清理连接池)
print("应用关闭,开始清理 ASR 连接池...")
# await close_asr_pool()
# 生命周期函数
@asynccontextmanager
async def lifespan(app: FastAPI):
print("=== 开始初始化 ASR 服务 ===")
init_success, init_msg = await ASRManager.initialize()
print(f"ASR 初始化结果:{init_msg}")
# if not init_success:
# 连接池为空/初始化失败,终止服务启动
# raise ServiceInitError(f"服务启动失败:{init_msg}")
# 2. 初始化WebSocketConnectionManager(绑定全局ConnectionManager
await ws_manager.initialize()
print("=== ASR 服务初始化完成 ===")
yield # 应用运行中
# 关闭时清理
print("=== 开始关闭 ASR 服务 ===")
await ASRManager.close()
print("=== ASR 服务关闭完成 ===")
app = FastAPI(
title="语音AI对话系统",
File diff suppressed because it is too large Load Diff
+238
View File
@@ -0,0 +1,238 @@
# audio_ai_chat/websocket/manager.py
from fastapi import WebSocket
from typing import Dict, Optional
from datetime import datetime
import base64
from audio_ai_chat.asr.base import ASRBase, ASRResultCallback
from audio_ai_chat.asr.asr_manager import ASRManager
from audio_ai_chat.websocket.connection_context import ConnectionManager, ConnectionContext # 导入全局单例类
from audio_ai_chat.config.logger import logger
class WebSocketConnectionManager:
"""全局唯一的WebSocket连接处理器(管理WebSocket连接生命周期)"""
def __init__(self):
self.asr_conn_map: Dict[str, Optional[object]] = {} # key=client_idvalue=ASR连接
# 不实例化新的ConnectionManager,而是使用全局单例
self.connection_manager: Optional[ConnectionManager] = None
async def initialize(self):
"""初始化:获取ConnectionManager全局单例(在FastAPI启动时调用)"""
self.connection_manager = await ConnectionManager.get_instance()
logger.info("WebSocketConnectionManager 初始化成功(绑定全局ConnectionManager")
async def handle_connection(self, websocket: WebSocket):
"""处理单个WebSocket连接的完整生命周期"""
# 校验ConnectionManager是否初始化
if not self.connection_manager:
await websocket.accept()
await websocket.send_text("服务未初始化完成,请稍后重试")
await websocket.close()
logger.error("WebSocketConnectionManager 未初始化,拒绝连接")
return
# 1. 接受连接,生成client_id(用字符串类型,避免int溢出)
await websocket.accept()
client_id = str(id(websocket)) # client_id为字符串,与ConnectionManager的key类型一致
logger.info(f"新WebSocket连接:client_id={client_id}")
try:
# 2. 创建连接上下文(通过全局ConnectionManager
context = await self.connection_manager.create_connection(client_id=client_id)
if not context:
await websocket.send_text("连接上下文创建失败")
await websocket.close()
return
# 3. 获取ASR实例
asr_client = ASRManager.get_instance()
if not asr_client or not ASRManager.is_available():
await websocket.send_json({
"type": "error",
"message": "ASR服务未初始化,无法提供转写服务",
"timestamp": datetime.utcnow().isoformat() + "Z"
})
await self.connection_manager.remove_connection(client_id=client_id)
await websocket.close()
return
# 4. 获取ASR连接
asr_conn = await asr_client.get_connection()
if not asr_conn:
await websocket.send_json({
"type": "error",
"message": "ASR无空闲连接,连接失败",
"timestamp": datetime.utcnow().isoformat() + "Z"
})
await self.connection_manager.remove_connection(client_id=client_id)
await websocket.close()
return
self.asr_conn_map[client_id] = asr_conn
# 5. 定义ASR结果回调(绑定当前上下文)
async def asr_callback(result: Dict[str, Any]):
if not context.is_active:
logger.warning(f"连接已关闭,忽略ASR结果:client_id={client_id}")
return
# 处理ASR结果并存入上下文
context.add_asr_result(result)
# 推送给前端
if result.get("error"):
await websocket.send_json({
"type": "asr_error",
"message": result["error"],
"timestamp": datetime.utcnow().isoformat() + "Z"
})
else:
await websocket.send_json({
"type": "asr_progress" if not result["is_final"] else "asr_final",
"text": result["text"],
"is_final": result["is_final"],
"timestamp": result.get("timestamp", datetime.utcnow().isoformat() + "Z")
})
# 6. 启动ASR通信
asr_task = asyncio.create_task(
asr_client.start_communication(conn=asr_conn, callback=asr_callback)
)
# 7. 循环接收前端数据
while context.is_active:
try:
# 假设前端发送JSON格式数据(区分音频/文本/用户信息)
data = await websocket.receive_json()
data_type = data.get("type")
# 处理用户信息(登录后发送)
if data_type == "user_info":
try:
token = data.get("token")
user_id = data.get("user_id")
name = data.get("name", "匿名用户")
context.set_user_info(token=token, user_id=user_id, name=name)
await websocket.send_json({
"type": "info",
"message": "用户信息设置成功",
"timestamp": datetime.utcnow().isoformat() + "Z"
})
except Exception as e:
err_msg = f"用户信息设置失败:{str(e)}"
context.add_system_message(err_msg)
await websocket.send_json({
"type": "error",
"message": err_msg,
"timestamp": datetime.utcnow().isoformat() + "Z"
})
# 处理Base64编码的音频数据
elif data_type == "audio_data":
audio_base64 = data.get("audio_data")
if not audio_base64:
continue
try:
audio_data = base64.b64decode(audio_base64)
success = await asr_client.push_audio(asr_conn, audio_data)
if not success:
await websocket.send_json({
"type": "warning",
"message": "ASR音频队列已满,部分数据丢失",
"timestamp": datetime.utcnow().isoformat() + "Z"
})
except Exception as e:
err_msg = f"音频解码失败:{str(e)}"
logger.error(f"client_id={client_id}{err_msg}")
await websocket.send_json({
"type": "error",
"message": err_msg,
"timestamp": datetime.utcnow().isoformat() + "Z"
})
# 处理纯文本输入
elif data_type == "text_input":
text = data.get("text", "").strip()
if text:
context.add_chat_history({
"role": "user",
"content": text,
"source": "text",
"asr_metadata": None
})
await websocket.send_json({
"type": "info",
"message": f"已接收文本:{text}",
"timestamp": datetime.utcnow().isoformat() + "Z"
})
# 处理大模型请求
elif data_type == "request_llm":
if context.is_processing:
await websocket.send_json({
"type": "warning",
"message": "正在处理上一个请求,请稍后再试",
"timestamp": datetime.utcnow().isoformat() + "Z"
})
continue
# 获取对话历史
chat_history = context.get_chat_history(limit=20)
logger.debug(f"请求大模型:client_id={client_id},历史条数={len(chat_history)}")
# 模拟大模型调用(实际替换为真实LLM调用)
context.is_processing = True
try:
# llm_response = await context.llm_session.generate(chat_history=chat_history)
llm_response = f"模拟大模型回复:已收到你的{len(chat_history)}条对话历史"
context.add_llm_result(llm_response)
await websocket.send_json({
"type": "llm_response",
"text": llm_response,
"timestamp": datetime.utcnow().isoformat() + "Z"
})
except Exception as e:
err_msg = f"大模型调用失败:{str(e)}"
context.add_system_message(err_msg)
await websocket.send_json({
"type": "error",
"message": err_msg,
"timestamp": datetime.utcnow().isoformat() + "Z"
})
finally:
context.is_processing = False
# 未知数据类型
else:
err_msg = f"未知数据类型:{data_type}"
await websocket.send_json({
"type": "error",
"message": err_msg,
"timestamp": datetime.utcnow().isoformat() + "Z"
})
except Exception as e:
# 捕获前端发送数据异常(如断开连接)
logger.error(f"接收前端数据异常:client_id={client_id}error={str(e)}")
break
except Exception as e:
# 其他异常
err_msg = f"连接处理异常:{str(e)}"
logger.error(f"client_id={client_id}{err_msg}")
await websocket.send_json({
"type": "error",
"message": err_msg,
"timestamp": datetime.utcnow().isoformat() + "Z"
})
finally:
# 8. 资源清理
# 取消ASR任务
asr_task.cancel()
try:
await asr_task
except asyncio.CancelledError:
pass
# 释放ASR连接
if client_id in self.asr_conn_map:
asr_conn = self.asr_conn_map.pop(client_id)
await asr_client.release_connection(asr_conn)
# 移除连接上下文
await self.connection_manager.remove_connection(client_id=client_id)
# 关闭WebSocket
await websocket.close()
logger.info(f"WebSocket连接关闭:client_id={client_id}")