1
This commit is contained in:
@@ -0,0 +1 @@
|
||||
https://github.com/modelscope/FunASR/blob/main/runtime/docs/SDK_advanced_guide_offline_gpu_zh.md
|
||||
@@ -0,0 +1,10 @@
|
||||
|
||||
git clone https://github.com/FunAudioLLM/Fun-ASR.git
|
||||
cd Fun-ASR
|
||||
pip install -r requirements.txt
|
||||
清华大学源
|
||||
|
||||
|
||||
pip install torch==2.5.1 torchaudio==2.5.1 --index-url https://download.pytorch.org/whl/cu121
|
||||
|
||||
pip install -r requirements.txt -i https://pypi.tuna.tsinghua.edu.cn/simple
|
||||
Binary file not shown.
@@ -22,7 +22,7 @@ class MessageType(IntEnum):
|
||||
CONTROL_CMD = 0b0100 # 控制指令
|
||||
IDENTITY = 0b0101 # 身份校验包json格式
|
||||
ERROR = 0b0110 # 错误信息json格式
|
||||
# 预留12种类型(0b0100 ~ 0b1111)
|
||||
TIP_MESSAGE = 0b0111 # 大模型生成的回答提示
|
||||
|
||||
|
||||
class SerializationType(IntEnum):
|
||||
@@ -40,12 +40,16 @@ class CompressionType(IntEnum):
|
||||
# 预留6种方式(0b011 ~ 0b111)
|
||||
|
||||
|
||||
class ControlCommand(IntEnum):
|
||||
class ControlCode(IntEnum):
|
||||
"""控制指令类型(配合MessageType.CONTROL_CMD使用)"""
|
||||
HEARTBEAT = 0b0001 # 心跳响应
|
||||
PAUSE = 0b0010 # 暂停
|
||||
RESUME = 0b0011 # 继续
|
||||
STOP = 0b0100 # 停止
|
||||
FINISH_PUSH = 0b0001, # 本轮tts推流结束
|
||||
FINISH_PLAY = 0b0010, # 语音播放结束(发后端)
|
||||
AI_ANSWER_BEGIN = 0b0011, # 本轮AI回答的文本内容开始
|
||||
AI_ANSWER_OVER = 0b0100, # 本轮AI回答的文本内容已经全部返回
|
||||
AI_CLUE = 0b0101, # 需要AI提示命令(发后端)
|
||||
AI_CLUE_OVER = 0b0110, # AI提示命令已经全部返回
|
||||
FRONT_TO_SERVER_OVER = 0b0111, # 语音按钮已经抬起
|
||||
SERVER_TO_FRONT_OVER = 0b1000, # 后端已经返回本次全部内容
|
||||
|
||||
|
||||
# -------------------------- 协议工具类(与JS协议结构一致) --------------------------
|
||||
@@ -89,7 +93,7 @@ class ProtocolCodec:
|
||||
if serialization is None and msg_type != MessageType.PING:
|
||||
if msg_type == MessageType.AUDIO_DATA:
|
||||
serialization = SerializationType.RAW
|
||||
elif msg_type == MessageType.TEXT_MESSAGE:
|
||||
elif msg_type == MessageType.TEXT_MESSAGE or msg_type == MessageType.TIP_MESSAGE:
|
||||
serialization = SerializationType.STRING
|
||||
elif msg_type == MessageType.CONTROL_CMD:
|
||||
serialization = SerializationType.JSON
|
||||
@@ -267,7 +271,7 @@ if __name__ == "__main__":
|
||||
print(f"文本解析结果:类型={msg_type2.name},顺序号={seq2},内容={body2}\n")
|
||||
|
||||
# 示例3:控制指令(JSON序列化)
|
||||
control_body = {"cmd": ControlCommand.PAUSE.value, "reason": "用户主动暂停"}
|
||||
control_body = {"cmd": ControlCode.PAUSE.value, "reason": "用户主动暂停"}
|
||||
control_packet = ProtocolCodec.pack(
|
||||
msg_type=MessageType.CONTROL_CMD,
|
||||
body=control_body,
|
||||
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -11,7 +11,7 @@ class ASRResult(TypedDict):
|
||||
"""ASR 识别结果的结构化定义(子类必须遵循此格式)"""
|
||||
client_id: str # 唯一请求 ID(用于关联 WebSocket 连接)
|
||||
text: str # 识别文本结果
|
||||
is_final: bool # 是否是最终结果(True:一句话结束,False:中间结果)
|
||||
is_final: bool # 是否是最终结果
|
||||
confidence: Optional[float] # 置信度(可选)
|
||||
error: Optional[str] # 错误信息(None 表示成功)
|
||||
timestamp: int # 结果生成时间戳(毫秒)
|
||||
@@ -101,7 +101,7 @@ class ASRBase(ABC):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def release_connection(self, conn: T, client_id: str) -> None:
|
||||
async def release_connection(self, conn: T) -> None:
|
||||
"""
|
||||
释放ASR连接(归还到连接池)
|
||||
参数:
|
||||
|
||||
Binary file not shown.
@@ -91,6 +91,7 @@ class FunASR(ASRBase, ABC):
|
||||
print("无效的ASR连接对象")
|
||||
return False
|
||||
if not conn or not conn.is_alive or conn.stop_event.is_set():
|
||||
print('111111', not conn)
|
||||
return False
|
||||
try:
|
||||
conn.audio_queue.put_nowait(audio_data)
|
||||
@@ -220,7 +221,7 @@ class FunASR(ASRBase, ABC):
|
||||
asr_conn.is_alive = False
|
||||
await result_callback(cast(ASRError, {
|
||||
"client_id": client_id,
|
||||
"error": f"音频发送失败:{str(e)}", # 注意:你之前的写法少了 f 字符串,这里修正
|
||||
"error": f"音频发送失败:{str(e)}",
|
||||
"code": 500,
|
||||
"timestamp": int(time.time() * 1000)
|
||||
}))
|
||||
@@ -233,13 +234,14 @@ class FunASR(ASRBase, ABC):
|
||||
try:
|
||||
asr_result = await asr_conn.ws.recv()
|
||||
result_json = json.loads(asr_result)
|
||||
if result_json.get("timestamp", "") == '':
|
||||
continue
|
||||
print('result_json', result_json)
|
||||
# if result_json.get("timestamp", "") == '':
|
||||
# continue
|
||||
result: ASRResult = {
|
||||
"client_id": client_id,
|
||||
"text": result_json.get("text", ""),
|
||||
"timestamp": result_json.get("timestamp", ""),
|
||||
"is_final": result_json.get("is_final", True),
|
||||
"timestamp": result_json.get("timestamp", None),
|
||||
"is_final": result_json.get("mode") == "2pass-offline",
|
||||
"error": None,
|
||||
"confidence": None
|
||||
}
|
||||
|
||||
@@ -0,0 +1,183 @@
|
||||
import asyncio
|
||||
import logging
|
||||
from typing import Dict, Optional
|
||||
from enum import Enum
|
||||
from fastapi import WebSocket, WebSocketDisconnect, FastAPI
|
||||
from audio_ai_chat.config.logger import logger
|
||||
from audio_ai_chat.core.asr.asr_manager import ASRManager
|
||||
from audio_ai_chat.codec.ProtocolCodec import ProtocolCodec, MessageType, ControlCode
|
||||
from audio_ai_chat.core.connection import ConnectionContext
|
||||
from audio_ai_chat.utils.exceptions import ServiceCallError
|
||||
from typing import Optional, Callable, Awaitable, Generic, TypeVar, TypedDict, Union
|
||||
from audio_ai_chat.core.asr.base import ASRResult, ASRError
|
||||
|
||||
# ------------------------------
|
||||
# 2. ASR服务管理器(独立ASR)
|
||||
# ------------------------------
|
||||
class ASRWebSocketManager:
|
||||
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self._send_task = None
|
||||
self._asr_task: Optional[asyncio.Task] = None
|
||||
self.is_active = False
|
||||
|
||||
# 通用资源清理
|
||||
async def _cleanup_resources(self):
|
||||
client_id = self.context.client_id if self.context else "未知"
|
||||
self.is_active = False
|
||||
logger.info(f"客户端 {client_id} 开始清理连接")
|
||||
|
||||
# 唤醒并等待发送协程退出
|
||||
if self._send_task and not self._send_task.done():
|
||||
try:
|
||||
self.context.return_ws_queue.put_nowait(b"")
|
||||
except asyncio.QueueFull:
|
||||
pass
|
||||
|
||||
if self._send_task:
|
||||
try:
|
||||
await asyncio.wait_for(self._send_task, timeout=2.0)
|
||||
except asyncio.TimeoutError:
|
||||
logger.warning(f"客户端 {client_id} 发送协程退出超时,强制取消")
|
||||
self._send_task.cancel()
|
||||
try:
|
||||
await self._send_task
|
||||
except asyncio.CancelledError:
|
||||
logger.info(f"客户端 {client_id} 发送协程已强制取消")
|
||||
|
||||
# await self._unregister_connection()
|
||||
self.context = None
|
||||
self._send_task = None
|
||||
logger.info(f"客户端 {client_id} 连接清理完成")
|
||||
|
||||
async def _send_worker(self):
|
||||
"""专用发送协程:支持状态判断、异常捕获、取消响应"""
|
||||
if not self.context:
|
||||
logger.warning("发送协程启动失败:ConnectionContext 未初始化")
|
||||
return
|
||||
|
||||
client_id = self.context.client_id
|
||||
websocket = self.context.websocket # 提前获取,避免重复访问
|
||||
logger.info(f"客户端 {client_id} 发送协程启动")
|
||||
|
||||
try:
|
||||
while self.is_active:
|
||||
try:
|
||||
# 超时时间可根据业务调整(建议0.1-1秒)
|
||||
message = await asyncio.wait_for(self.context.return_ws_queue.get(), timeout=0.05)
|
||||
except asyncio.TimeoutError:
|
||||
# 超时后重新进入循环,检查 is_active 和连接状态
|
||||
# if not self.is_active or websocket.client_state != 1:
|
||||
# break
|
||||
continue # 继续等待消息
|
||||
|
||||
# 4. 安全发送:捕获所有可能的异常
|
||||
try:
|
||||
await websocket.send_bytes(message)
|
||||
logger.debug(f"客户端 {client_id} 发送消息:{len(message)} 字节")
|
||||
except WebSocket.Disconnect:
|
||||
logger.info(f"客户端 {client_id} 已断开,发送失败(连接失效)")
|
||||
except Exception as e:
|
||||
logger.error(f"客户端 {client_id} 消息发送异常:{e}", exc_info=True)
|
||||
finally:
|
||||
# 5. 必须标记任务完成(避免任务泄漏)
|
||||
self.context.return_ws_queue.task_done()
|
||||
|
||||
# 6. 退出前清理:批量处理剩余消息(无需发送,仅标记完成)
|
||||
# remaining = self.context.return_ws_queue.qsize()
|
||||
# if remaining > 0:
|
||||
# logger.info(f"客户端 {client_id} 发送协程退出,清理剩余 {remaining} 条消息")
|
||||
# while not self.context.return_ws_queue.empty():
|
||||
# self.context.return_ws_queue.get_nowait()
|
||||
# self.context.return_ws_queue.task_done()
|
||||
|
||||
except asyncio.CancelledError:
|
||||
# 7. 捕获协程取消异常(正常退出,无需报错)
|
||||
logger.info(f"客户端 {client_id} 发送协程被强制取消")
|
||||
except Exception as e:
|
||||
logger.error(f"客户端 {client_id} 发送协程异常退出:{e}", exc_info=True)
|
||||
finally:
|
||||
logger.info(f"客户端 {client_id} 发送协程已退出")
|
||||
|
||||
async def _asr_callback(self, result: Union[ASRResult, ASRError]) -> None:
|
||||
"""ASR识别结果回调:发送给客户端"""
|
||||
await asyncio.sleep(0)
|
||||
if self.is_active and self.context:
|
||||
text = result.get("text")
|
||||
is_final = result.get("is_final")
|
||||
print(text, is_final)
|
||||
if is_final:
|
||||
result_packet = ProtocolCodec.pack(MessageType.TEXT_MESSAGE, text)
|
||||
else:
|
||||
result_packet = ProtocolCodec.pack(MessageType.TIP_MESSAGE, text)
|
||||
await self.context.return_ws_queue.put(result_packet)
|
||||
|
||||
async def handle_asr_connection(self, websocket: WebSocket):
|
||||
client_id = str(id(websocket))
|
||||
asr_client = None
|
||||
asr_conn = None
|
||||
try:
|
||||
# 建立连接
|
||||
await websocket.accept()
|
||||
logger.info(f"ASR客户端 {client_id} 已连接")
|
||||
self.context = ConnectionContext(client_id=client_id)
|
||||
self.context.websocket = websocket
|
||||
# 初始化ASR(仅ASR服务需要)
|
||||
# 获取asr连接实例
|
||||
asr_client = ASRManager.get_instance()
|
||||
asr_conn = await asr_client.get_connection(client_id)
|
||||
self.is_active = True
|
||||
if not asr_conn:
|
||||
raise ServiceCallError("获取ASR连接失败")
|
||||
|
||||
# 启动ASR通信协程
|
||||
self._asr_task = asyncio.create_task(
|
||||
asr_client.start_communication(asr_conn, self._asr_callback, client_id)
|
||||
)
|
||||
# 启动发送协程
|
||||
self._send_task = asyncio.create_task(self._send_worker())
|
||||
logger.info(f"ASR客户端 {client_id} 就绪")
|
||||
|
||||
# 处理音频数据
|
||||
while self.is_active:
|
||||
raw_bytes = await websocket.receive_bytes()
|
||||
logger.debug(f"收到ASR客户端 {client_id} 音频数据:{len(raw_bytes)} 字节")
|
||||
msg_type, sequence, data = ProtocolCodec.unpack(raw_bytes)
|
||||
|
||||
if msg_type == MessageType.AUDIO_DATA:
|
||||
# 推送音频到ASR
|
||||
success = await asr_client.push_audio(asr_conn, data, client_id)
|
||||
if not success:
|
||||
logger.warning(f"ASR客户端 {client_id} 音频推送失败")
|
||||
elif msg_type == MessageType.CONTROL_CMD:
|
||||
# 处理ASR控制指令
|
||||
print('处理ASR控制指令')
|
||||
if data.get("type") == ControlCode.FRONT_TO_SERVER_OVER: # 抬手指令
|
||||
result_packet = ProtocolCodec.pack(MessageType.CONTROL_CMD, {"type": ControlCode.SERVER_TO_FRONT_OVER})
|
||||
await self.context.return_ws_queue.put(result_packet)
|
||||
|
||||
|
||||
else:
|
||||
logger.warning(f"ASR客户端 {client_id} 收到未知消息类型:{msg_type.value}")
|
||||
|
||||
except WebSocketDisconnect:
|
||||
logger.info(f"ASR客户端 {client_id} 主动断开连接")
|
||||
except Exception as e:
|
||||
logger.error(f"ASR客户端 {client_id} 处理异常:{e}", exc_info=True)
|
||||
finally:
|
||||
if asr_client:
|
||||
await asr_client.release_connection(asr_conn)
|
||||
# 额外清理ASR协程
|
||||
if self._asr_task and not self._asr_task.done():
|
||||
self._asr_task.cancel()
|
||||
try:
|
||||
await self._asr_task
|
||||
except asyncio.CancelledError:
|
||||
logger.info(f"ASR客户端 {client_id} ASR协程已取消")
|
||||
await self._cleanup_resources()
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -89,17 +89,6 @@ class WebSocketConnectionManager:
|
||||
# break
|
||||
continue # 继续等待消息
|
||||
|
||||
# # 2. 核心退出判断:连接已断开或协程被标记为非活跃
|
||||
# if not self.is_active or websocket.client_state != 1:
|
||||
# logger.debug(f"客户端 {client_id} 连接已断开,放弃发送消息")
|
||||
# self.context.return_ws_queue.task_done()
|
||||
# break # 退出循环,协程结束
|
||||
#
|
||||
# # 3. 过滤唤醒消息(空消息是退出信号,无需发送)
|
||||
# if message == b"":
|
||||
# self.context.return_ws_queue.task_done()
|
||||
# continue
|
||||
|
||||
# 4. 安全发送:捕获所有可能的异常
|
||||
try:
|
||||
await websocket.send_bytes(message)
|
||||
@@ -128,73 +117,19 @@ class WebSocketConnectionManager:
|
||||
finally:
|
||||
logger.info(f"客户端 {client_id} 发送协程已退出")
|
||||
|
||||
async def connect(self, client_id: str, websocket: WebSocket) -> ConnectionContext:
|
||||
"""建立连接+身份校验"""
|
||||
await websocket.accept()
|
||||
logger.info(
|
||||
|
||||
f"连接 {client_id} 已接受,等待身份信息(5秒超时),当前连接数: {await self.get_active_count()}"
|
||||
)
|
||||
# 超时接收身份包
|
||||
try:
|
||||
ping_packet = await asyncio.wait_for(websocket.receive_bytes(), timeout=5.0)
|
||||
except asyncio.TimeoutError:
|
||||
error_msg = f"连接 {client_id} 身份校验超时"
|
||||
logger.warning(error_msg)
|
||||
error_packet = ProtocolCodec.pack(
|
||||
MessageType.ERROR, {"code": 1008, "message": "身份校验超时,请重试"}
|
||||
)
|
||||
await websocket.send_bytes(error_packet)
|
||||
raise TimeoutError(error_msg)
|
||||
|
||||
# 解包并验证包类型
|
||||
msg_type, _, identity_data = ProtocolCodec.unpack(ping_packet)
|
||||
# if msg_type != MessageType.IDENTITY:
|
||||
# error_msg = f"连接 {client_id} 首个包类型错误"
|
||||
# logger.error(error_msg)
|
||||
# error_packet = ProtocolCodec.pack(
|
||||
# MessageType.ERROR, {"code": 4001, "message": "非法请求:首个包必须是身份校验包"}
|
||||
# )
|
||||
# await websocket.send_bytes(error_packet)
|
||||
# raise ValueError(error_msg)
|
||||
|
||||
# 校验身份信息
|
||||
# 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} 身份信息不完整"
|
||||
# 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)
|
||||
|
||||
# 创建/获取连接上下文
|
||||
context = ConnectionContext(client_id=client_id)
|
||||
context.websocket = websocket
|
||||
# context.set_user_info(token, user_id, name)
|
||||
# 响应身份校验成功
|
||||
# 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
|
||||
|
||||
# ------------------------------
|
||||
# 核心连接处理逻辑(修复启动时机、优化退出流程)
|
||||
# 核心连接处理逻辑
|
||||
# ------------------------------
|
||||
async def handle_connection(self, websocket: WebSocket):
|
||||
client_id = str(id(websocket)) # 用websocket实例ID作为client_id(唯一)
|
||||
try:
|
||||
self.context = await self.connect(client_id, websocket)
|
||||
await websocket.accept()
|
||||
logger.info(f"连接 {client_id} 已接受,当前连接数: {await self.get_active_count()}")
|
||||
|
||||
# 创建连接上下文(无鉴权逻辑,直接初始化)
|
||||
self.context = ConnectionContext(client_id=client_id)
|
||||
self.context.websocket = websocket
|
||||
|
||||
self.is_active = True
|
||||
await self._register_connection(client_id, self.context) # 注册并且存储上下文
|
||||
|
||||
|
||||
@@ -9,7 +9,7 @@ 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
|
||||
|
||||
from audio_ai_chat.core.asr_websocket_manager import ASRWebSocketManager
|
||||
|
||||
|
||||
|
||||
@@ -54,7 +54,11 @@ async def websocket_audio(websocket: WebSocket):
|
||||
ws_manager = WebSocketConnectionManager()
|
||||
await ws_manager.handle_connection(websocket)
|
||||
|
||||
|
||||
@app.websocket("/ws/asr")
|
||||
async def websocket_asr(websocket: WebSocket):
|
||||
"""ASR语音识别服务WebSocket端点"""
|
||||
asr_manager = ASRWebSocketManager()
|
||||
await asr_manager.handle_asr_connection(websocket)
|
||||
|
||||
@app.get("/health")
|
||||
async def health_check():
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,31 @@
|
||||
# 下载 Miniconda(适配 Linux x86_64,其他架构可换链接)
|
||||
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh
|
||||
|
||||
# 执行安装脚本(一路回车,最后选 yes 初始化 conda)
|
||||
bash Miniconda3-latest-Linux-x86_64.sh
|
||||
|
||||
# 重启终端,或执行以下命令让 conda 生效
|
||||
source ~/.bashrc # 如果是 zsh 终端,用 source ~/.zshrc
|
||||
|
||||
|
||||
# 接受 main 频道的条款
|
||||
conda tos accept --override-channels --channel https://repo.anaconda.com/pkgs/main
|
||||
|
||||
# 接受 r 频道的条款
|
||||
conda tos accept --override-channels --channel https://repo.anaconda.com/pkgs/r
|
||||
|
||||
|
||||
# 创建并激活 FunASR 虚拟环境
|
||||
conda create -n funasr python=3.8 -y # -y 自动确认安装
|
||||
conda activate funasr # 激活后终端前缀会显示 (funasr)
|
||||
|
||||
本地目录虚拟环境
|
||||
conda create --prefix .venv python=3.10 -y
|
||||
激活
|
||||
conda activate ./.venv
|
||||
退出
|
||||
conda deactivate
|
||||
# 删除环境(直接删目录,或用 Conda 命令)
|
||||
rm -rf .venv # 最简单
|
||||
# 查看环境信息(确认路径)
|
||||
conda info --envs # 会显示 .venv 的绝对路径
|
||||
Reference in New Issue
Block a user