From c9bf6c501c3c26db33f31ca260ff65612be78ddf Mon Sep 17 00:00:00 2001 From: caixypromise Date: Sun, 28 Dec 2025 23:13:58 +0800 Subject: [PATCH] =?UTF-8?q?feat(sharedASR):=20=E4=BC=98=E5=8C=96ASR?= =?UTF-8?q?=E5=90=AF=E5=8A=A8=E7=AD=96=E7=95=A5=EF=BC=8C=E6=94=AF=E6=8C=81?= =?UTF-8?q?=E6=9C=AC=E5=9C=B0=E6=A8=A1=E5=9E=8B=E9=A2=84=E5=8A=A0=E8=BD=BD?= =?UTF-8?q?=EF=BC=8C=E9=81=BF=E5=85=8D=E8=BF=9E=E6=8E=A5=E8=B6=85=E6=97=B6?= =?UTF-8?q?=E3=80=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- main/xiaozhi-server/config.yaml | 5 +- .../config/mqtt_config_example.yaml | 54 +++ .../core/components/adapters/asr_adapter.py | 65 ++- .../core/processors/text_processor.py | 1 + .../core/protocols/mqtt_protocol.py | 1 + .../core/providers/asr/shared_asr_manager.py | 407 ++++++++++++++++++ .../core/providers/asr/shared_asr_proxy.py | 114 +++++ .../core/servers/mqtt_server.py | 1 + .../core/servers/multi_protocol_server.py | 1 + .../core/services/connection_service.py | 6 + .../core/transport/mqtt_transport.py | 1 + .../core/xiaozhi_server_facade.py | 95 +++- 12 files changed, 724 insertions(+), 27 deletions(-) create mode 100644 main/xiaozhi-server/config/mqtt_config_example.yaml create mode 100644 main/xiaozhi-server/core/providers/asr/shared_asr_manager.py create mode 100644 main/xiaozhi-server/core/providers/asr/shared_asr_proxy.py diff --git a/main/xiaozhi-server/config.yaml b/main/xiaozhi-server/config.yaml index 69c9e47f..e0767926 100644 --- a/main/xiaozhi-server/config.yaml +++ b/main/xiaozhi-server/config.yaml @@ -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 diff --git a/main/xiaozhi-server/config/mqtt_config_example.yaml b/main/xiaozhi-server/config/mqtt_config_example.yaml new file mode 100644 index 00000000..1d828044 --- /dev/null +++ b/main/xiaozhi-server/config/mqtt_config_example.yaml @@ -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发送音频数据 + diff --git a/main/xiaozhi-server/core/components/adapters/asr_adapter.py b/main/xiaozhi-server/core/components/adapters/asr_adapter.py index 776e81b6..34e32961 100644 --- a/main/xiaozhi-server/core/components/adapters/asr_adapter.py +++ b/main/xiaozhi-server/core/components/adapters/asr_adapter.py @@ -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): diff --git a/main/xiaozhi-server/core/processors/text_processor.py b/main/xiaozhi-server/core/processors/text_processor.py index d0d5e390..e0cc97fc 100644 --- a/main/xiaozhi-server/core/processors/text_processor.py +++ b/main/xiaozhi-server/core/processors/text_processor.py @@ -52,3 +52,4 @@ class TextProcessor(MessageProcessor): })) except Exception as e: logger.error(f"发送错误响应失败: {e}") + diff --git a/main/xiaozhi-server/core/protocols/mqtt_protocol.py b/main/xiaozhi-server/core/protocols/mqtt_protocol.py index f04e83b8..53d9b800 100644 --- a/main/xiaozhi-server/core/protocols/mqtt_protocol.py +++ b/main/xiaozhi-server/core/protocols/mqtt_protocol.py @@ -430,3 +430,4 @@ class MQTTProtocol: self.socket.close() except Exception as e: logger.error(f"关闭socket失败: {e}") + diff --git a/main/xiaozhi-server/core/providers/asr/shared_asr_manager.py b/main/xiaozhi-server/core/providers/asr/shared_asr_manager.py new file mode 100644 index 00000000..63cb9ae4 --- /dev/null +++ b/main/xiaozhi-server/core/providers/asr/shared_asr_manager.py @@ -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 + diff --git a/main/xiaozhi-server/core/providers/asr/shared_asr_proxy.py b/main/xiaozhi-server/core/providers/asr/shared_asr_proxy.py new file mode 100644 index 00000000..dca07610 --- /dev/null +++ b/main/xiaozhi-server/core/providers/asr/shared_asr_proxy.py @@ -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 + diff --git a/main/xiaozhi-server/core/servers/mqtt_server.py b/main/xiaozhi-server/core/servers/mqtt_server.py index d5b901e4..eb02d000 100644 --- a/main/xiaozhi-server/core/servers/mqtt_server.py +++ b/main/xiaozhi-server/core/servers/mqtt_server.py @@ -311,3 +311,4 @@ class UDPProtocol(asyncio.DatagramProtocol): def error_received(self, exc): logger.error(f"UDP协议错误: {exc}") + diff --git a/main/xiaozhi-server/core/servers/multi_protocol_server.py b/main/xiaozhi-server/core/servers/multi_protocol_server.py index 57ffd4fd..251e4baf 100644 --- a/main/xiaozhi-server/core/servers/multi_protocol_server.py +++ b/main/xiaozhi-server/core/servers/multi_protocol_server.py @@ -316,3 +316,4 @@ class MultiProtocolServer: def is_protocol_enabled(self, protocol: str) -> bool: """检查协议是否启用""" return protocol in self.servers + diff --git a/main/xiaozhi-server/core/services/connection_service.py b/main/xiaozhi-server/core/services/connection_service.py index d1f4ea09..e4204cf5 100644 --- a/main/xiaozhi-server/core/services/connection_service.py +++ b/main/xiaozhi-server/core/services/connection_service.py @@ -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 diff --git a/main/xiaozhi-server/core/transport/mqtt_transport.py b/main/xiaozhi-server/core/transport/mqtt_transport.py index 05eacb5c..4a3c24ba 100644 --- a/main/xiaozhi-server/core/transport/mqtt_transport.py +++ b/main/xiaozhi-server/core/transport/mqtt_transport.py @@ -243,3 +243,4 @@ class UDPAudioHandler: self._closed = True self.message_callback = None self.remote_address = None + diff --git a/main/xiaozhi-server/core/xiaozhi_server_facade.py b/main/xiaozhi-server/core/xiaozhi_server_facade.py index 40d84516..8579e251 100644 --- a/main/xiaozhi-server/core/xiaozhi_server_facade.py +++ b/main/xiaozhi-server/core/xiaozhi_server_facade.py @@ -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() } +