x
This commit is contained in:
+4
-3
@@ -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
@@ -0,0 +1,7 @@
|
||||
<component name="ProjectDictionaryState">
|
||||
<dictionary name="project">
|
||||
<words>
|
||||
<w>tymas</w>
|
||||
</words>
|
||||
</dictionary>
|
||||
</component>
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -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)
|
||||
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
Binary file not shown.
Binary file not shown.
@@ -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)
|
||||
@@ -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_id(int),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
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -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
|
||||
Binary file not shown.
Binary file not shown.
@@ -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()
|
||||
@@ -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
|
||||
Binary file not shown.
Binary file not shown.
@@ -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} 资源清理完成")
|
||||
@@ -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对话系统",
|
||||
|
||||
Binary file not shown.
Binary file not shown.
File diff suppressed because it is too large
Load Diff
@@ -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_id,value=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}")
|
||||
Reference in New Issue
Block a user