mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-23 07:33:53 +08:00
feat(sharedASR): 优化ASR启动策略,支持本地模型预加载,避免连接超时。
This commit is contained in:
@@ -297,12 +297,15 @@ Memory:
|
||||
# 如果这里不填,则会默认使用selected_module.LLM的模型作为意图识别的思考模型
|
||||
# 如果你的不想使用selected_module.LLM记忆存储,这里最好使用独立的LLM作为意图识别,例如使用免费的ChatGLMLLM
|
||||
llm: ChatGLMLLM
|
||||
|
||||
ASR:
|
||||
FunASR:
|
||||
type: fun_local
|
||||
model_dir: models/SenseVoiceSmall
|
||||
output_dir: tmp/
|
||||
# 队列最大大小(可选,默认100)
|
||||
# 当并发请求超过此值时,会返回"服务繁忙"提示
|
||||
# 建议根据服务器性能调整,GPU 服务器可适当增大
|
||||
queue_max_size: 100
|
||||
FunASRServer:
|
||||
# 独立部署FunASR,使用FunASR的API服务,只需要五句话
|
||||
# 第一句:mkdir -p ./funasr-runtime-resources/models
|
||||
|
||||
@@ -0,0 +1,54 @@
|
||||
# MQTT协议配置示例
|
||||
# 将此配置添加到你的主配置文件中
|
||||
|
||||
# 协议配置
|
||||
protocols:
|
||||
enabled_protocols: ["websocket", "mqtt"] # 启用的协议列表
|
||||
websocket_enabled: true # WebSocket协议开关
|
||||
mqtt_enabled: true # MQTT协议开关
|
||||
|
||||
# MQTT服务器配置
|
||||
mqtt_server:
|
||||
enabled: true # 是否启用MQTT服务器
|
||||
host: "0.0.0.0" # 监听地址
|
||||
port: 1883 # MQTT端口
|
||||
udp_port: 1883 # UDP端口(用于音频传输)
|
||||
public_ip: "your.server.ip" # 公网IP地址
|
||||
max_connections: 1000 # 最大连接数
|
||||
heartbeat_interval: 30 # 心跳检查间隔(秒)
|
||||
max_payload_size: 8192 # 最大消息载荷大小
|
||||
|
||||
# 服务器配置(扩展)
|
||||
server:
|
||||
ip: "0.0.0.0"
|
||||
port: 8080 # WebSocket端口
|
||||
http_port: 8081
|
||||
auth_key: ""
|
||||
vision_explain: ""
|
||||
|
||||
# MQTT服务器配置(嵌套)
|
||||
mqtt_server:
|
||||
enabled: true
|
||||
host: "0.0.0.0"
|
||||
port: 1883
|
||||
udp_port: 1883
|
||||
public_ip: "localhost"
|
||||
max_connections: 1000
|
||||
heartbeat_interval: 30
|
||||
max_payload_size: 8192
|
||||
|
||||
# 使用示例:
|
||||
# 1. 启动多协议服务器:
|
||||
# python main_multi_protocol.py
|
||||
#
|
||||
# 2. WebSocket客户端连接:
|
||||
# ws://your.server.ip:8080/
|
||||
#
|
||||
# 3. MQTT客户端连接:
|
||||
# mqtt://your.server.ip:1883
|
||||
# 客户端ID格式:GID_test@@@mac_address@@@uuid
|
||||
# 或:GID_test@@@mac_address
|
||||
#
|
||||
# 4. UDP音频传输:
|
||||
# 客户端通过MQTT接收UDP配置后,使用UDP发送音频数据
|
||||
|
||||
@@ -5,15 +5,23 @@ from core.utils.modules_initialize import initialize_asr
|
||||
from config.logger import setup_logging
|
||||
|
||||
logger = setup_logging()
|
||||
TAG = __name__
|
||||
|
||||
|
||||
class ASRAdapter(Component):
|
||||
"""ASR组件适配器:将现有ASR组件包装为新的组件接口"""
|
||||
"""
|
||||
ASR组件适配器:将现有ASR组件包装为新的组件接口
|
||||
|
||||
支持两种模式:
|
||||
1. 共享实例模式:使用 SharedASRManager 的全局共享实例
|
||||
2. 独立实例模式:每个连接创建独立的 ASR 实例(原有逻辑)
|
||||
"""
|
||||
|
||||
def __init__(self, config: Dict[str, Any]):
|
||||
super().__init__(ComponentType.ASR, config)
|
||||
self._asr_instance = None
|
||||
self._delete_audio = config.get("delete_audio", True)
|
||||
self._using_shared = False # 是否使用共享实例
|
||||
|
||||
async def _do_initialize(self, context: Any) -> None:
|
||||
"""初始化ASR组件"""
|
||||
@@ -23,40 +31,61 @@ class ASRAdapter(Component):
|
||||
if not selected_module:
|
||||
raise ValueError("未配置ASR模块")
|
||||
|
||||
# 创建ASR实例
|
||||
self._asr_instance = initialize_asr(self.config)
|
||||
# 检查是否有全局共享 ASR 管理器
|
||||
shared_manager = getattr(context, 'shared_asr_manager', None)
|
||||
|
||||
# 注册资源以便清理
|
||||
self.add_resource(self._asr_instance)
|
||||
if shared_manager and shared_manager.is_ready():
|
||||
# 使用共享实例模式
|
||||
logger.bind(tag=TAG).info(f"使用共享 ASR 实例: {selected_module}")
|
||||
from core.providers.asr.shared_asr_proxy import SharedASRProxy
|
||||
self._asr_instance = SharedASRProxy(shared_manager)
|
||||
self._using_shared = True
|
||||
else:
|
||||
# 使用独立实例模式(原有逻辑)
|
||||
logger.bind(tag=TAG).info(f"使用独立 ASR 实例: {selected_module}")
|
||||
self._asr_instance = initialize_asr(self.config)
|
||||
self._using_shared = False
|
||||
|
||||
# 打开音频通道
|
||||
# 注册资源以便清理(仅非共享实例)
|
||||
if not self._using_shared:
|
||||
self.add_resource(self._asr_instance)
|
||||
|
||||
# 打开音频通道(如果需要)
|
||||
if hasattr(self._asr_instance, 'open_audio_channels'):
|
||||
await self._asr_instance.open_audio_channels(context)
|
||||
|
||||
logger.info(f"ASR组件初始化完成: {selected_module}")
|
||||
logger.bind(tag=TAG).info(
|
||||
f"ASR组件初始化完成: {selected_module}, "
|
||||
f"共享模式: {self._using_shared}"
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"ASR组件初始化失败: {e}")
|
||||
logger.bind(tag=TAG).error(f"ASR组件初始化失败: {e}")
|
||||
raise
|
||||
|
||||
async def _do_cleanup(self) -> None:
|
||||
"""清理ASR组件"""
|
||||
if self._asr_instance:
|
||||
try:
|
||||
# 关闭ASR实例
|
||||
if hasattr(self._asr_instance, 'close'):
|
||||
await self._asr_instance.close()
|
||||
|
||||
# 清理音频文件
|
||||
if hasattr(self._asr_instance, 'cleanup_audio_files'):
|
||||
self._asr_instance.cleanup_audio_files()
|
||||
|
||||
logger.info("ASR组件清理完成")
|
||||
# 如果是共享实例,不需要关闭(由服务器统一管理)
|
||||
if self._using_shared:
|
||||
logger.bind(tag=TAG).debug("共享 ASR 实例,跳过清理")
|
||||
else:
|
||||
# 关闭独立 ASR 实例
|
||||
if hasattr(self._asr_instance, 'close'):
|
||||
await self._asr_instance.close()
|
||||
|
||||
# 清理音频文件
|
||||
if hasattr(self._asr_instance, 'cleanup_audio_files'):
|
||||
self._asr_instance.cleanup_audio_files()
|
||||
|
||||
logger.bind(tag=TAG).info("ASR组件清理完成")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"ASR组件清理失败: {e}")
|
||||
logger.bind(tag=TAG).error(f"ASR组件清理失败: {e}")
|
||||
finally:
|
||||
self._asr_instance = None
|
||||
self._using_shared = False
|
||||
|
||||
@property
|
||||
def asr_instance(self):
|
||||
|
||||
@@ -52,3 +52,4 @@ class TextProcessor(MessageProcessor):
|
||||
}))
|
||||
except Exception as e:
|
||||
logger.error(f"发送错误响应失败: {e}")
|
||||
|
||||
|
||||
@@ -430,3 +430,4 @@ class MQTTProtocol:
|
||||
self.socket.close()
|
||||
except Exception as e:
|
||||
logger.error(f"关闭socket失败: {e}")
|
||||
|
||||
|
||||
@@ -0,0 +1,407 @@
|
||||
"""
|
||||
SharedASRManager: 全局 ASR 管理器
|
||||
实现单例模型 + 单推理执行器 + 队列限流。
|
||||
单例的原因是:推理是 CPU/GPU-bound,不是 I/O-bound,多实例不仅会占用内存,还会降低吞吐能力
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Dict, Any, Optional, Tuple, List
|
||||
from config.logger import setup_logging
|
||||
|
||||
logger = setup_logging()
|
||||
TAG = __name__
|
||||
|
||||
|
||||
class SharedASRManager:
|
||||
"""
|
||||
全局共享 ASR 管理器
|
||||
"""
|
||||
|
||||
# 支持预加载的本地模型类型
|
||||
LOCAL_MODEL_TYPES = [
|
||||
"fun_local", # FunASR 本地
|
||||
"sherpa_onnx_local", # Sherpa ONNX
|
||||
"sense_voice" # SenseVoice
|
||||
]
|
||||
|
||||
def __init__(self, config: Dict[str, Any], asr_type: str = None):
|
||||
"""
|
||||
初始化 ASR 管理器
|
||||
Args:
|
||||
config: 服务器配置
|
||||
asr_type: ASR 类型(Optional,用于显式指定)
|
||||
"""
|
||||
self.config = config
|
||||
self.asr_type = asr_type
|
||||
|
||||
# 模型实例(全局单例)
|
||||
self.model_instance = None
|
||||
|
||||
# 任务队列(限流)
|
||||
queue_max_size = self._get_queue_max_size()
|
||||
self.task_queue: asyncio.Queue = asyncio.Queue(maxsize=queue_max_size)
|
||||
|
||||
# 推理锁(使得推理串行化)
|
||||
self.inference_lock = asyncio.Lock()
|
||||
# 线程池执行器,用于阻塞调用
|
||||
self.executor: Optional[ThreadPoolExecutor] = None
|
||||
|
||||
# 运行状态
|
||||
self.running = False
|
||||
self._inference_task: Optional[asyncio.Task] = None
|
||||
self.is_local_model = self._check_local_model()
|
||||
|
||||
logger.bind(tag=TAG).info(
|
||||
f"SharedASRManager 初始化完成, "
|
||||
f"类型: {self.asr_type}, "
|
||||
f"本地模型: {self.is_local_model}, "
|
||||
f"队列大小: {queue_max_size}"
|
||||
)
|
||||
|
||||
def _get_queue_max_size(self) -> int:
|
||||
"""获取队列最大大小"""
|
||||
# 尝试从配置获取
|
||||
selected_asr = self.config.get("selected_module", {}).get("ASR")
|
||||
if selected_asr:
|
||||
asr_config = self.config.get("ASR", {}).get(selected_asr, {})
|
||||
return asr_config.get("queue_max_size", 100)
|
||||
return 100
|
||||
|
||||
def _check_local_model(self) -> bool:
|
||||
"""检查是否为本地模型"""
|
||||
if self.asr_type:
|
||||
return self.asr_type in self.LOCAL_MODEL_TYPES
|
||||
|
||||
# 从配置推断
|
||||
selected_asr = self.config.get("selected_module", {}).get("ASR")
|
||||
if not selected_asr:
|
||||
return False
|
||||
|
||||
asr_config = self.config.get("ASR", {}).get(selected_asr, {})
|
||||
asr_type = asr_config.get("type", selected_asr)
|
||||
self.asr_type = asr_type
|
||||
|
||||
return asr_type in self.LOCAL_MODEL_TYPES
|
||||
|
||||
async def initialize(self):
|
||||
"""
|
||||
初始化管理器
|
||||
- 预加载模型
|
||||
- 启动推理执行器
|
||||
"""
|
||||
if not self.is_local_model:
|
||||
logger.bind(tag=TAG).info("非本地模型,跳过预加载")
|
||||
return
|
||||
|
||||
if self.running:
|
||||
logger.bind(tag=TAG).warning("管理器已在运行中")
|
||||
return
|
||||
|
||||
try:
|
||||
logger.bind(tag=TAG).info(f"开始预加载 ASR 模型: {self.asr_type}")
|
||||
|
||||
# 预加载模型
|
||||
await self._preload_model()
|
||||
|
||||
# 启动推理执行器
|
||||
self.running = True
|
||||
self._inference_task = asyncio.create_task(self._inference_loop())
|
||||
|
||||
logger.bind(tag=TAG).info("ASR 模型预加载完成,推理执行器已启动")
|
||||
|
||||
except Exception as e:
|
||||
logger.bind(tag=TAG).error(f"ASR 模型预加载失败: {e}")
|
||||
raise
|
||||
|
||||
async def _preload_model(self):
|
||||
"""在线程池中预加载模型"""
|
||||
# 创建线程池
|
||||
self.executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix="asr_worker")
|
||||
|
||||
loop = asyncio.get_event_loop()
|
||||
self.model_instance = await loop.run_in_executor(
|
||||
self.executor,
|
||||
self._create_model_instance
|
||||
)
|
||||
|
||||
logger.bind(tag=TAG).info("模型实例创建完成")
|
||||
|
||||
def _create_model_instance(self):
|
||||
"""
|
||||
实际创建模型实例(在线程中执行)
|
||||
|
||||
Returns:
|
||||
ASR Provider 实例
|
||||
"""
|
||||
from core.utils.modules_initialize import initialize_asr
|
||||
|
||||
logger.bind(tag=TAG).info("正在创建 ASR 模型实例...")
|
||||
instance = initialize_asr(self.config)
|
||||
logger.bind(tag=TAG).info("ASR 模型实例创建成功")
|
||||
|
||||
return instance
|
||||
|
||||
async def submit_task(
|
||||
self,
|
||||
opus_data: List[bytes],
|
||||
session_id: str,
|
||||
audio_format: str = "opus"
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""
|
||||
提交推理任务
|
||||
|
||||
Args:
|
||||
opus_data: 音频数据
|
||||
session_id: 会话 ID
|
||||
audio_format: 音频格式
|
||||
|
||||
Returns:
|
||||
(识别文本, 文件路径)
|
||||
|
||||
Raises:
|
||||
RuntimeError: 队列满或服务未运行
|
||||
"""
|
||||
if not self.running:
|
||||
raise RuntimeError("ASR 服务未运行")
|
||||
|
||||
# 检查队列是否满(限流)
|
||||
if self.task_queue.full():
|
||||
queue_status = self.get_queue_status()
|
||||
logger.bind(tag=TAG).warning(
|
||||
f"ASR 队列已满: {queue_status}"
|
||||
)
|
||||
raise RuntimeError("ASR 服务繁忙,请稍后重试")
|
||||
|
||||
# 创建 Future 用于返回结果
|
||||
result_future: asyncio.Future = asyncio.Future()
|
||||
|
||||
# 构造任务
|
||||
task = {
|
||||
'opus_data': opus_data,
|
||||
'session_id': session_id,
|
||||
'audio_format': audio_format,
|
||||
'future': result_future
|
||||
}
|
||||
|
||||
# 放入队列
|
||||
await self.task_queue.put(task)
|
||||
|
||||
logger.bind(tag=TAG).debug(
|
||||
f"任务已提交, session: {session_id}, "
|
||||
f"队列大小: {self.task_queue.qsize()}"
|
||||
)
|
||||
|
||||
# 等待结果
|
||||
return await result_future
|
||||
|
||||
async def _inference_loop(self):
|
||||
"""
|
||||
单个推理执行器循环
|
||||
|
||||
核心原则:
|
||||
- 只有一个执行器
|
||||
- 串行处理任务
|
||||
- 带超时的队列获取,支持优雅退出
|
||||
"""
|
||||
logger.bind(tag=TAG).info("推理执行器启动")
|
||||
|
||||
while self.running:
|
||||
task = None
|
||||
try:
|
||||
# 带超时的队列获取,避免关闭时卡住
|
||||
try:
|
||||
task = await asyncio.wait_for(
|
||||
self.task_queue.get(),
|
||||
timeout=1.0
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
# 超时后检查 running 状态,继续循环
|
||||
continue
|
||||
|
||||
# 执行推理(加锁保证串行)
|
||||
async with self.inference_lock:
|
||||
result = await self._run_inference(
|
||||
task['opus_data'],
|
||||
task['session_id'],
|
||||
task['audio_format']
|
||||
)
|
||||
|
||||
# 返回结果
|
||||
if not task['future'].done():
|
||||
task['future'].set_result(result)
|
||||
|
||||
logger.bind(tag=TAG).debug(
|
||||
f"推理完成, session: {task['session_id']}"
|
||||
)
|
||||
|
||||
except asyncio.CancelledError:
|
||||
logger.bind(tag=TAG).info("推理执行器被取消")
|
||||
break
|
||||
except Exception as e:
|
||||
logger.bind(tag=TAG).error(f"推理执行失败: {e}")
|
||||
if task and 'future' in task and not task['future'].done():
|
||||
task['future'].set_exception(e)
|
||||
|
||||
logger.bind(tag=TAG).info("推理执行器已停止")
|
||||
|
||||
async def _run_inference(
|
||||
self,
|
||||
opus_data: List[bytes],
|
||||
session_id: str,
|
||||
audio_format: str
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""
|
||||
执行实际推理(在线程池中)
|
||||
|
||||
Args:
|
||||
opus_data: 音频数据
|
||||
session_id: 会话 ID
|
||||
audio_format: 音频格式
|
||||
|
||||
Returns:
|
||||
(识别文本, 文件路径)
|
||||
"""
|
||||
loop = asyncio.get_event_loop()
|
||||
|
||||
# 在线程池中执行推理
|
||||
result = await loop.run_in_executor(
|
||||
self.executor,
|
||||
lambda: self.model_instance.speech_to_text_sync(
|
||||
opus_data, session_id, audio_format
|
||||
) if hasattr(self.model_instance, 'speech_to_text_sync')
|
||||
else self._sync_wrapper(opus_data, session_id, audio_format)
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
def _sync_wrapper(
|
||||
self,
|
||||
opus_data: List[bytes],
|
||||
session_id: str,
|
||||
audio_format: str
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""
|
||||
同步包装器
|
||||
处理 async speech_to_text 方法
|
||||
"""
|
||||
import asyncio
|
||||
|
||||
async def _call():
|
||||
return await self.model_instance.speech_to_text(
|
||||
opus_data, session_id, audio_format
|
||||
)
|
||||
|
||||
# 创建新的事件循环执行
|
||||
loop = None
|
||||
try:
|
||||
loop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
result = loop.run_until_complete(_call())
|
||||
return result
|
||||
finally:
|
||||
if loop:
|
||||
loop.close()
|
||||
|
||||
async def shutdown(self):
|
||||
"""
|
||||
优雅停机
|
||||
|
||||
步骤:
|
||||
1. 停止接收新任务
|
||||
2. 等待当前任务完成(带超时)
|
||||
3. 取消未完成的任务
|
||||
4. 关闭线程池
|
||||
"""
|
||||
if not self.running:
|
||||
return
|
||||
|
||||
logger.bind(tag=TAG).info("开始关闭 ASR 管理器...")
|
||||
|
||||
# 停止接收新任务
|
||||
self.running = False
|
||||
|
||||
# 等待推理任务完成
|
||||
if self._inference_task and not self._inference_task.done():
|
||||
try:
|
||||
# 最多等待 5 秒
|
||||
await asyncio.wait_for(
|
||||
self._inference_task,
|
||||
timeout=5.0
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.bind(tag=TAG).warning("推理任务超时,强制取消")
|
||||
self._inference_task.cancel()
|
||||
try:
|
||||
await self._inference_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
# 取消所有队列中未完成的任务
|
||||
cancelled_count = 0
|
||||
while not self.task_queue.empty():
|
||||
try:
|
||||
task = self.task_queue.get_nowait()
|
||||
if not task['future'].done():
|
||||
task['future'].set_exception(
|
||||
RuntimeError("ASR 服务正在关闭")
|
||||
)
|
||||
cancelled_count += 1
|
||||
except asyncio.QueueEmpty:
|
||||
break
|
||||
|
||||
if cancelled_count > 0:
|
||||
logger.bind(tag=TAG).info(f"已取消 {cancelled_count} 个待处理任务")
|
||||
|
||||
# 关闭线程池
|
||||
if self.executor:
|
||||
self.executor.shutdown(wait=False)
|
||||
self.executor = None
|
||||
logger.bind(tag=TAG).info("线程池已关闭")
|
||||
|
||||
# 清理模型实例
|
||||
self.model_instance = None
|
||||
|
||||
logger.bind(tag=TAG).info("ASR 管理器已关闭")
|
||||
|
||||
def get_queue_status(self) -> Dict[str, Any]:
|
||||
"""
|
||||
获取队列状态(用于监控)
|
||||
|
||||
Returns:
|
||||
队列状态字典
|
||||
"""
|
||||
queue_size = self.task_queue.qsize()
|
||||
queue_max = self.task_queue.maxsize
|
||||
|
||||
return {
|
||||
'queue_size': queue_size,
|
||||
'queue_max': queue_max,
|
||||
'is_busy': queue_size > queue_max * 0.8,
|
||||
'utilization': queue_size / queue_max if queue_max > 0 else 0,
|
||||
'running': self.running
|
||||
}
|
||||
|
||||
def is_ready(self) -> bool:
|
||||
"""检查管理器是否就绪"""
|
||||
return (
|
||||
self.running and
|
||||
self.model_instance is not None and
|
||||
self.executor is not None
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def is_local_model_type(cls, asr_type: str) -> bool:
|
||||
"""
|
||||
检查 ASR 类型是否为本地模型
|
||||
|
||||
Args:
|
||||
asr_type: ASR 类型
|
||||
|
||||
Returns:
|
||||
是否为本地模型
|
||||
"""
|
||||
return asr_type in cls.LOCAL_MODEL_TYPES
|
||||
|
||||
@@ -0,0 +1,114 @@
|
||||
"""
|
||||
SharedASRProxy: 共享 ASR 管理器的代理类
|
||||
|
||||
功能:
|
||||
- 包装 SharedASRManager
|
||||
- 提供与原 ASR Provider 相同的接口
|
||||
- 处理队列满等异常情况
|
||||
"""
|
||||
|
||||
from typing import List, Tuple, Optional, Dict, Any
|
||||
from core.providers.asr.base import ASRProviderBase
|
||||
from core.providers.asr.dto.dto import InterfaceType
|
||||
from config.logger import setup_logging
|
||||
|
||||
logger = setup_logging()
|
||||
TAG = __name__
|
||||
|
||||
|
||||
class SharedASRProxy(ASRProviderBase):
|
||||
"""
|
||||
共享 ASR 管理器的代理类
|
||||
|
||||
该类提供与原 ASR Provider 相同的接口,
|
||||
但实际推理工作由 SharedASRManager 完成。
|
||||
"""
|
||||
|
||||
def __init__(self, manager):
|
||||
"""
|
||||
初始化代理
|
||||
|
||||
Args:
|
||||
manager: SharedASRManager 实例
|
||||
"""
|
||||
super().__init__()
|
||||
self.manager = manager
|
||||
|
||||
# 从共享管理器获取接口类型
|
||||
if manager.model_instance and hasattr(manager.model_instance, 'interface_type'):
|
||||
self.interface_type = manager.model_instance.interface_type
|
||||
else:
|
||||
self.interface_type = InterfaceType.LOCAL
|
||||
|
||||
logger.bind(tag=TAG).info("SharedASRProxy 初始化完成")
|
||||
|
||||
async def speech_to_text(
|
||||
self,
|
||||
opus_data: List[bytes],
|
||||
session_id: str,
|
||||
audio_format: str = "opus"
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""
|
||||
语音转文本(通过共享管理器)
|
||||
|
||||
Args:
|
||||
opus_data: 音频数据(Opus 编码的字节列表)
|
||||
session_id: 会话 ID
|
||||
audio_format: 音频格式,默认 "opus"
|
||||
|
||||
Returns:
|
||||
Tuple[str, str]: (识别的文本, 音频文件路径)
|
||||
"""
|
||||
try:
|
||||
# 检查管理器状态
|
||||
if not self.manager.is_ready():
|
||||
logger.bind(tag=TAG).error("ASR 管理器未就绪")
|
||||
return "", None
|
||||
|
||||
# 提交任务到共享管理器
|
||||
result = await self.manager.submit_task(
|
||||
opus_data,
|
||||
session_id,
|
||||
audio_format
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
except RuntimeError as e:
|
||||
# 队列满或服务未运行
|
||||
logger.bind(tag=TAG).warning(f"ASR 服务繁忙: {e}")
|
||||
# 返回友好提示,而不是空字符串
|
||||
return "服务繁忙,请稍后重试", None
|
||||
|
||||
except Exception as e:
|
||||
logger.bind(tag=TAG).error(f"ASR 推理失败: {e}")
|
||||
return "", None
|
||||
|
||||
def get_queue_status(self) -> Dict[str, Any]:
|
||||
"""
|
||||
获取队列状态
|
||||
|
||||
Returns:
|
||||
队列状态字典
|
||||
"""
|
||||
return self.manager.get_queue_status()
|
||||
|
||||
def is_ready(self) -> bool:
|
||||
"""
|
||||
检查代理是否就绪
|
||||
|
||||
Returns:
|
||||
是否就绪
|
||||
"""
|
||||
return self.manager.is_ready()
|
||||
|
||||
async def close(self):
|
||||
"""
|
||||
关闭代理
|
||||
|
||||
注意:不关闭共享管理器,由服务器统一管理
|
||||
"""
|
||||
logger.bind(tag=TAG).debug("SharedASRProxy 关闭")
|
||||
# 代理不负责关闭共享管理器
|
||||
pass
|
||||
|
||||
@@ -311,3 +311,4 @@ class UDPProtocol(asyncio.DatagramProtocol):
|
||||
|
||||
def error_received(self, exc):
|
||||
logger.error(f"UDP协议错误: {exc}")
|
||||
|
||||
|
||||
@@ -316,3 +316,4 @@ class MultiProtocolServer:
|
||||
def is_protocol_enabled(self, protocol: str) -> bool:
|
||||
"""检查协议是否启用"""
|
||||
return protocol in self.servers
|
||||
|
||||
|
||||
@@ -42,6 +42,12 @@ class ConnectionService:
|
||||
# 设置transport接口
|
||||
context.transport = transport
|
||||
|
||||
# 传入共享 ASR 管理器(如果有)
|
||||
# 这使得 ASRAdapter 可以使用预加载的模型实例
|
||||
if '_shared_asr_manager' in self.config:
|
||||
context.shared_asr_manager = self.config['_shared_asr_manager']
|
||||
logger.debug("连接使用共享 ASR 实例")
|
||||
|
||||
# 兼容性:设置websocket属性(如果transport是WebSocket)
|
||||
if hasattr(transport, '_websocket'):
|
||||
context.websocket = transport._websocket
|
||||
|
||||
@@ -243,3 +243,4 @@ class UDPAudioHandler:
|
||||
self._closed = True
|
||||
self.message_callback = None
|
||||
self.remote_address = None
|
||||
|
||||
|
||||
@@ -11,12 +11,18 @@ from config.config_loader import get_protocol_config, get_mqtt_server_config
|
||||
from core.servers.multi_protocol_server import MultiProtocolServer
|
||||
|
||||
logger = setup_logging()
|
||||
TAG = __name__
|
||||
|
||||
|
||||
class XiaozhiServerFacade:
|
||||
"""
|
||||
小智服务器门面类
|
||||
提供统一的服务器管理接口,屏蔽内部协议复杂性
|
||||
|
||||
功能:
|
||||
- 协议管理(WebSocket、MQTT)
|
||||
- 本地 ASR 模型预加载
|
||||
- 优雅启动和停止
|
||||
"""
|
||||
|
||||
def __init__(self, config: Dict[str, Any]):
|
||||
@@ -28,6 +34,7 @@ class XiaozhiServerFacade:
|
||||
"""
|
||||
self.config = config
|
||||
self.multi_protocol_server: Optional[MultiProtocolServer] = None
|
||||
self.shared_asr_manager = None # 共享 ASR 管理器
|
||||
self.is_initialized = False
|
||||
self.is_running = False
|
||||
|
||||
@@ -116,22 +123,73 @@ class XiaozhiServerFacade:
|
||||
async def initialize(self):
|
||||
"""初始化服务器"""
|
||||
if self.is_initialized:
|
||||
logger.warning("服务器已经初始化")
|
||||
logger.bind(tag=TAG).warning("服务器已经初始化")
|
||||
return
|
||||
|
||||
try:
|
||||
logger.info("正在初始化小智服务器...")
|
||||
logger.bind(tag=TAG).info("正在初始化小智服务器...")
|
||||
|
||||
# 检查并预加载本地 ASR 模型(关键步骤)
|
||||
await self._preload_asr_if_needed()
|
||||
|
||||
# 创建多协议服务器
|
||||
self.multi_protocol_server = MultiProtocolServer(self.config)
|
||||
|
||||
self.is_initialized = True
|
||||
logger.info("小智服务器初始化完成")
|
||||
logger.bind(tag=TAG).info("小智服务器初始化完成")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"初始化服务器失败: {e}")
|
||||
logger.bind(tag=TAG).error(f"初始化服务器失败: {e}")
|
||||
raise
|
||||
|
||||
async def _preload_asr_if_needed(self):
|
||||
"""
|
||||
检查并预加载本地 ASR 模型
|
||||
|
||||
如果配置使用本地 ASR 模型(如 FunASR),则在服务器启动时预加载,
|
||||
避免首次语音识别时的延迟导致客户端超时。
|
||||
"""
|
||||
try:
|
||||
# 获取 ASR 配置
|
||||
selected_asr = self.config.get("selected_module", {}).get("ASR")
|
||||
if not selected_asr:
|
||||
logger.bind(tag=TAG).info("未配置 ASR 模块,跳过预加载")
|
||||
return
|
||||
|
||||
# 获取 ASR 类型
|
||||
asr_config = self.config.get("ASR", {}).get(selected_asr, {})
|
||||
asr_type = asr_config.get("type", selected_asr)
|
||||
|
||||
# 导入 SharedASRManager 检查是否为本地模型
|
||||
from core.providers.asr.shared_asr_manager import SharedASRManager
|
||||
|
||||
if SharedASRManager.is_local_model_type(asr_type):
|
||||
logger.bind(tag=TAG).info(
|
||||
f"检测到本地 ASR 模型: {asr_type},开始预加载..."
|
||||
)
|
||||
|
||||
# 创建全局 ASR 管理器
|
||||
self.shared_asr_manager = SharedASRManager(self.config, asr_type)
|
||||
|
||||
# 预加载模型
|
||||
await self.shared_asr_manager.initialize()
|
||||
|
||||
# 将管理器放入配置中供后续使用
|
||||
self.config['_shared_asr_manager'] = self.shared_asr_manager
|
||||
|
||||
logger.bind(tag=TAG).info(
|
||||
f"ASR 模型预加载完成,类型: {asr_type}"
|
||||
)
|
||||
else:
|
||||
logger.bind(tag=TAG).info(
|
||||
f"ASR 类型为远程服务: {asr_type},无需预加载"
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.bind(tag=TAG).error(f"ASR 预加载失败: {e}")
|
||||
# 预加载失败不影响服务器启动,继续使用懒加载模式
|
||||
logger.bind(tag=TAG).warning("将回退到懒加载模式")
|
||||
|
||||
async def start(self):
|
||||
"""启动服务器"""
|
||||
if not self.is_initialized:
|
||||
@@ -158,20 +216,30 @@ class XiaozhiServerFacade:
|
||||
async def stop(self):
|
||||
"""停止服务器"""
|
||||
if not self.is_running:
|
||||
logger.info("服务器未在运行")
|
||||
logger.bind(tag=TAG).info("服务器未在运行")
|
||||
return
|
||||
|
||||
try:
|
||||
logger.info("正在停止小智服务器...")
|
||||
logger.bind(tag=TAG).info("正在停止小智服务器...")
|
||||
|
||||
# 停止多协议服务器
|
||||
if self.multi_protocol_server:
|
||||
await self.multi_protocol_server.stop()
|
||||
|
||||
# 关闭共享 ASR 管理器(优雅停机)
|
||||
if self.shared_asr_manager:
|
||||
logger.bind(tag=TAG).info("正在关闭共享 ASR 管理器...")
|
||||
await self.shared_asr_manager.shutdown()
|
||||
self.shared_asr_manager = None
|
||||
# 从配置中移除
|
||||
if '_shared_asr_manager' in self.config:
|
||||
del self.config['_shared_asr_manager']
|
||||
|
||||
self.is_running = False
|
||||
logger.info("小智服务器已停止")
|
||||
logger.bind(tag=TAG).info("小智服务器已停止")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"停止服务器失败: {e}")
|
||||
logger.bind(tag=TAG).error(f"停止服务器失败: {e}")
|
||||
|
||||
async def restart(self):
|
||||
"""重启服务器"""
|
||||
@@ -225,6 +293,16 @@ class XiaozhiServerFacade:
|
||||
server_status = self.multi_protocol_server.get_server_status()
|
||||
base_status.update(server_status)
|
||||
|
||||
# 添加 ASR 状态
|
||||
if self.shared_asr_manager:
|
||||
base_status['asr'] = {
|
||||
'mode': 'shared',
|
||||
'ready': self.shared_asr_manager.is_ready(),
|
||||
'queue_status': self.shared_asr_manager.get_queue_status()
|
||||
}
|
||||
else:
|
||||
base_status['asr'] = {'mode': 'lazy_load'}
|
||||
|
||||
return base_status
|
||||
|
||||
def get_active_connections_count(self) -> Dict[str, int]:
|
||||
@@ -289,3 +367,4 @@ class XiaozhiServerFacade:
|
||||
'mqtt': self.get_mqtt_info(),
|
||||
'active_connections': self.get_active_connections_count()
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user