diff --git a/main/xiaozhi-server/app.py b/main/xiaozhi-server/app.py
index 0edb475b..cbb8f2e7 100644
--- a/main/xiaozhi-server/app.py
+++ b/main/xiaozhi-server/app.py
@@ -3,11 +3,11 @@ import uuid
import signal
import asyncio
from aioconsole import ainput
-from config.settings import load_config
+from config.config_loader import load_config
from config.logger import setup_logging
from core.utils.util import get_local_ip, validate_mcp_endpoint
from core.http_server import SimpleHttpServer
-from core.websocket_server import WebSocketServer
+from core.xiaozhi_server_facade import XiaozhiServerFacade
from core.utils.util import check_ffmpeg_installed
TAG = __name__
@@ -57,9 +57,10 @@ async def main():
# 添加 stdin 监控任务
stdin_task = asyncio.create_task(monitor_stdin())
- # 启动 WebSocket 服务器
- ws_server = WebSocketServer(config)
- ws_task = asyncio.create_task(ws_server.start())
+ # 启动小智服务器门面(支持WebSocket和MQTT)
+ xiaozhi_server = XiaozhiServerFacade(config)
+ xiaozhi_task = asyncio.create_task(xiaozhi_server.start())
+
# 启动 Simple http 服务器
ota_server = SimpleHttpServer(config)
ota_task = asyncio.create_task(ota_server.start())
@@ -89,24 +90,52 @@ async def main():
logger.bind(tag=TAG).error("mcp接入点不符合规范")
config["mcp_endpoint"] = "你的接入点 websocket地址"
- # 获取WebSocket配置,使用安全的默认值
- websocket_port = 8000
- server_config = config.get("server", {})
- if isinstance(server_config, dict):
- websocket_port = int(server_config.get("port", 8000))
-
- logger.bind(tag=TAG).info(
- "Websocket地址是\tws://{}:{}/xiaozhi/v1/",
- get_local_ip(),
- websocket_port,
- )
-
- logger.bind(tag=TAG).info(
- "=======上面的地址是websocket协议地址,请勿用浏览器访问======="
- )
- logger.bind(tag=TAG).info(
- "如想测试websocket请用谷歌浏览器打开test目录下的test_page.html"
- )
+ # 显示协议连接信息
+ connection_info = xiaozhi_server.get_connection_info()
+
+ # WebSocket信息
+ websocket_info = connection_info.get('websocket', {})
+ if websocket_info.get('enabled', False):
+ websocket_port = websocket_info.get('port', 8000)
+ logger.bind(tag=TAG).info(
+ "WebSocket地址是\tws://{}:{}/xiaozhi/v1/",
+ get_local_ip(),
+ websocket_port,
+ )
+
+ # MQTT信息
+ mqtt_info = connection_info.get('mqtt', {})
+ if mqtt_info.get('enabled', False):
+ mqtt_port = mqtt_info.get('port', 1883)
+ udp_port = mqtt_info.get('udp_port', 1883)
+ logger.bind(tag=TAG).info(
+ "MQTT地址是\t\tmqtt://{}:{}",
+ get_local_ip(),
+ mqtt_port,
+ )
+ logger.bind(tag=TAG).info(
+ "UDP音频端口是\t{}:{}",
+ get_local_ip(),
+ udp_port,
+ )
+
+ # 显示启用的协议
+ enabled_protocols = xiaozhi_server.config.get('enabled_protocols', [])
+ logger.bind(tag=TAG).info(f"启用的协议: {', '.join(enabled_protocols)}")
+
+ if 'websocket' in enabled_protocols:
+ logger.bind(tag=TAG).info(
+ "=======上面的WebSocket地址请勿用浏览器访问======="
+ )
+ logger.bind(tag=TAG).info(
+ "如想测试WebSocket请用谷歌浏览器打开test目录下的test_page.html"
+ )
+
+ if 'mqtt' in enabled_protocols:
+ logger.bind(tag=TAG).info(
+ "=======MQTT客户端ID格式: GID_test@@@mac_address@@@uuid======="
+ )
+
logger.bind(tag=TAG).info(
"=============================================================\n"
)
@@ -116,15 +145,25 @@ async def main():
except asyncio.CancelledError:
print("任务被取消,清理资源中...")
finally:
+ # 停止小智服务器
+ try:
+ await xiaozhi_server.stop()
+ except Exception as e:
+ logger.error(f"停止小智服务器失败: {e}")
+
# 取消所有任务(关键修复点)
stdin_task.cancel()
- ws_task.cancel()
+ xiaozhi_task.cancel()
if ota_task:
ota_task.cancel()
# 等待任务终止(必须加超时)
+ tasks_to_wait = [stdin_task, xiaozhi_task]
+ if ota_task:
+ tasks_to_wait.append(ota_task)
+
await asyncio.wait(
- [stdin_task, ws_task, ota_task] if ota_task else [stdin_task, ws_task],
+ tasks_to_wait,
timeout=3.0,
return_when=asyncio.ALL_COMPLETED,
)
diff --git a/main/xiaozhi-server/config.yaml b/main/xiaozhi-server/config.yaml
index c0aa3f93..69c9e47f 100644
--- a/main/xiaozhi-server/config.yaml
+++ b/main/xiaozhi-server/config.yaml
@@ -41,6 +41,54 @@ server:
# 可选:设备白名单,如果设置了白名单,那么白名单的机器无论是什么token都可以连接。
#allowed_devices:
# - "24:0A:C4:1D:3B:F0" # MAC地址列表
+
+# #####################################################################################
+# #############################协议配置(Protocol Configuration)########################
+# 支持WebSocket和MQTT两种协议,可以单独启用或同时启用
+protocols:
+ # 启用的协议列表,可选值: ["websocket", "mqtt"]
+ enabled_protocols: ["websocket"] # 默认只启用WebSocket
+ # WebSocket协议开关
+ websocket_enabled: true
+ # MQTT协议开关
+ mqtt_enabled: false
+
+# MQTT服务器配置(仅在mqtt_enabled为true时生效)
+mqtt_server:
+ # 是否启用MQTT服务器
+ enabled: false
+ # MQTT服务器监听地址
+ host: 0.0.0.0
+ # MQTT服务器端口
+ port: 1883
+ # UDP音频传输端口(通常与MQTT端口相同)
+ udp_port: 1883
+ # 公网IP地址(用于UDP音频传输配置)
+ # 如果使用docker部署或公网部署,请设置为实际的公网IP或域名
+ public_ip: localhost
+ # 最大连接数
+ max_connections: 1000
+ # 心跳检查间隔(秒)
+ heartbeat_interval: 30
+ # 最大消息载荷大小(字节)
+ max_payload_size: 8192
+
+# MQTT协议使用说明:
+# 1. 客户端ID格式:GID_test@@@mac_address@@@uuid 或 GID_test@@@mac_address
+# 例如:GID_test@@@aa:bb:cc:dd:ee:ff@@@unique_uuid_123
+# 2. 连接地址:mqtt://your.server.ip:1883
+# 3. 音频传输:通过UDP加密传输,配置信息在hello消息中返回
+# 4. 消息格式:JSON格式,支持hello、音频、文本等消息类型
+#
+# 启用MQTT的配置示例:
+# protocols:
+# enabled_protocols: ["websocket", "mqtt"] # 同时启用两种协议
+# mqtt_enabled: true
+# mqtt_server:
+# enabled: true
+# port: 1883
+# public_ip: "your.server.ip" # 替换为实际IP
+
log:
# 设置控制台输出的日志格式,时间、日志级别、标签、消息
log_format: "{time:YYMMDD HH:mm:ss}[{version}_{selected_module}][{extra[tag]}]-{level}-{message}"
diff --git a/main/xiaozhi-server/config/config_loader.py b/main/xiaozhi-server/config/config_loader.py
index 4f35b1fd..1e08b01c 100644
--- a/main/xiaozhi-server/config/config_loader.py
+++ b/main/xiaozhi-server/config/config_loader.py
@@ -1,8 +1,229 @@
import os
import yaml
from collections.abc import Mapping
+from typing import Any, Dict, Optional, Type, TypeVar, Union, get_type_hints, get_origin, get_args
+from dataclasses import dataclass, field, fields, MISSING
+import inspect
from config.manage_api_client import init_service, get_server_config, get_agent_models
+T = TypeVar('T')
+
+
+class ConfigDict(dict):
+ """增强的配置字典,支持点号访问和嵌套获取"""
+
+ def __init__(self, data: Dict[str, Any] = None):
+ super().__init__()
+ if data:
+ for key, value in data.items():
+ if isinstance(value, dict):
+ self[key] = ConfigDict(value)
+ else:
+ self[key] = value
+
+ def __getattr__(self, key: str) -> Any:
+ """支持点号访问"""
+ try:
+ return self[key]
+ except KeyError:
+ raise AttributeError(f"'{self.__class__.__name__}' object has no attribute '{key}'")
+
+ def __setattr__(self, key: str, value: Any) -> None:
+ """支持点号设置"""
+ self[key] = value
+
+ def __getitem__(self, key: str) -> Any:
+ """重写[]访问,支持嵌套路径,找不到抛出KeyError"""
+ if '.' in key:
+ keys = key.split('.')
+ current = self
+ for k in keys:
+ if not isinstance(current, (dict, ConfigDict)):
+ raise KeyError(f"Cannot access '{k}' on non-dict object")
+ current = super(ConfigDict, current).__getitem__(k)
+ return current
+ return super().__getitem__(key)
+
+ def get(self, key: str, default: Any = None) -> Any:
+ """重写get方法,支持嵌套路径"""
+ try:
+ return self[key]
+ except KeyError:
+ return default
+
+ def __setitem__(self, key: str, value: Any) -> None:
+ """重写[]设置,支持嵌套路径"""
+ if '.' in key:
+ keys = key.split('.')
+ current = self
+ for k in keys[:-1]:
+ if k not in current:
+ current[k] = ConfigDict()
+ elif not isinstance(current[k], (dict, ConfigDict)):
+ current[k] = ConfigDict()
+ current = current[k]
+
+ if isinstance(value, dict) and not isinstance(value, ConfigDict):
+ value = ConfigDict(value)
+ super(ConfigDict, current).__setitem__(keys[-1], value)
+ else:
+ if isinstance(value, dict) and not isinstance(value, ConfigDict):
+ value = ConfigDict(value)
+ super().__setitem__(key, value)
+
+
+class ConfigField:
+ """配置字段,模仿dataclass的field功能"""
+
+ def __init__(self, default=MISSING, default_factory=MISSING, prefix: str = None):
+ self.default = default
+ self.default_factory = default_factory
+ self.prefix = prefix
+
+ if default is not MISSING and default_factory is not MISSING:
+ raise ValueError("Cannot specify both default and default_factory")
+
+
+def config_field(default=MISSING, default_factory=MISSING, prefix: str = None):
+ """创建配置字段"""
+ return ConfigField(default, default_factory, prefix)
+
+
+def _is_config_class(cls: Type) -> bool:
+ """检查类是否是ConfigurationProperties装饰的配置类"""
+ return hasattr(cls, '_config_prefix') and hasattr(cls, '_inject_config')
+
+
+def _create_nested_config_instance(config_class: Type, config: ConfigDict, config_path: str):
+ """创建嵌套配置类实例"""
+ try:
+ # 获取嵌套配置数据 - 直接传递整个config,让嵌套类自己处理前缀
+ # 因为嵌套类有自己的prefix,它会从config中正确提取数据
+ return config_class(config)
+ except Exception as e:
+ # 如果创建失败,返回None或抛出更详细的错误
+ raise ValueError(f"Failed to create nested config instance for {config_class.__name__} at path '{config_path}': {e}")
+
+
+def ConfigurationProperties(prefix: str = "", auto_inject: bool = True):
+ """
+ 配置属性装饰器,模仿Spring Boot的@ConfigurationProperties
+
+ Args:
+ prefix: 配置前缀,如 'server.database'
+ auto_inject: 是否自动注入配置
+ """
+ def decorator(cls: Type[T]) -> Type[T]:
+ if not inspect.isclass(cls):
+ raise TypeError("ConfigurationProperties can only be applied to classes")
+
+ # 保存原始的__init__方法
+ original_init = cls.__init__ if hasattr(cls, '__init__') else None
+
+ # 获取类的类型注解
+ type_hints = get_type_hints(cls)
+
+ def new_init(self, config: ConfigDict = None, **kwargs):
+ # 如果有原始的__init__,先调用它
+ if original_init and original_init is not object.__init__:
+ try:
+ original_init(self)
+ except TypeError:
+ # 如果原始__init__不接受参数,忽略
+ pass
+
+ if config is None:
+ # 如果没有传入config,尝试从全局获取
+ config = getattr(self.__class__, '_global_config', None)
+ if config is None:
+ return
+
+ # 注入配置
+ self._inject_config(config, prefix, **kwargs)
+
+ def _inject_config(self, config: ConfigDict, config_prefix: str = "", **overrides):
+ """注入配置到实例属性"""
+ # 处理类属性
+ for attr_name in dir(self.__class__):
+ if attr_name.startswith('_'):
+ continue
+
+ attr_value = getattr(self.__class__, attr_name)
+ if isinstance(attr_value, ConfigField):
+ # 确定配置路径
+ field_prefix = attr_value.prefix or config_prefix
+ config_path = f"{field_prefix}.{attr_name}" if field_prefix else attr_name
+
+ # 从overrides或config获取值
+ if attr_name in overrides:
+ value = overrides[attr_name]
+ else:
+ # 检查是否有类型注解,如果是嵌套配置类则特殊处理
+ attr_type = type_hints.get(attr_name)
+ if attr_type and inspect.isclass(attr_type) and _is_config_class(attr_type):
+ # 嵌套配置类,创建实例
+ try:
+ value = _create_nested_config_instance(attr_type, config, config_path)
+ except ValueError as e:
+ # 如果创建失败,使用默认值
+ print(f"Warning: {e}")
+ if attr_value.default_factory is not MISSING:
+ value = attr_value.default_factory()
+ else:
+ value = attr_value.default
+ else:
+ # 普通类型,使用默认值逻辑
+ if attr_value.default_factory is not MISSING:
+ default_val = attr_value.default_factory()
+ else:
+ default_val = attr_value.default
+
+ value = config.get(config_path, default_val)
+
+ setattr(self, attr_name, value)
+
+ # 处理类型注解的属性
+ for attr_name, attr_type in type_hints.items():
+ if hasattr(self, attr_name):
+ continue # 已经通过ConfigField处理过了
+
+ config_path = f"{config_prefix}.{attr_name}" if config_prefix else attr_name
+
+ if attr_name in overrides:
+ value = overrides[attr_name]
+ else:
+ # 检查是否是嵌套的ConfigurationProperties类
+ if inspect.isclass(attr_type) and _is_config_class(attr_type):
+ # 创建嵌套配置类实例
+ try:
+ value = _create_nested_config_instance(attr_type, config, config_path)
+ except ValueError as e:
+ # 如果创建失败,使用None或默认值
+ print(f"Warning: {e}")
+ value = None
+ else:
+ # 普通类型,直接从配置获取
+ value = config.get(config_path)
+
+ if value is not None:
+ setattr(self, attr_name, value)
+
+ # 添加方法到类
+ cls.__init__ = new_init
+ cls._inject_config = _inject_config
+ cls._config_prefix = prefix
+
+ # 添加类方法用于设置全局配置
+ @classmethod
+ def set_global_config(cls, config: ConfigDict):
+ cls._global_config = config
+
+ cls.set_global_config = set_global_config
+
+ return cls
+
+ return decorator
+
def get_project_dir():
"""获取项目根目录"""
@@ -22,6 +243,9 @@ def load_config():
# 检查缓存
cached_config = cache_manager.get(CacheType.CONFIG, "main_config")
if cached_config is not None:
+ # 确保返回的是ConfigDict类型
+ if not isinstance(cached_config, ConfigDict):
+ cached_config = ConfigDict(cached_config)
return cached_config
default_config_path = get_project_dir() + "config.yaml"
@@ -36,6 +260,10 @@ def load_config():
else:
# 合并配置
config = merge_configs(default_config, custom_config)
+
+ # 转换为ConfigDict
+ config = ConfigDict(config)
+
# 初始化目录
ensure_directories(config)
@@ -68,7 +296,7 @@ def get_config_from_api(config):
"vision_explain": config["server"].get("vision_explain", ""),
"auth_key": config["server"].get("auth_key", ""),
}
- return config_data
+ return ConfigDict(config_data)
def get_private_config_from_api(config, device_id, client_id):
@@ -146,3 +374,218 @@ def merge_configs(default_config, custom_config):
merged[key] = value
return merged
+
+
+# 导出主要的类和函数
+__all__ = [
+ 'ConfigDict',
+ 'ConfigField',
+ 'config_field',
+ 'ConfigurationProperties',
+ 'load_config',
+ 'get_project_dir',
+ 'merge_configs'
+]
+
+
+# 配置类定义
+
+@ConfigurationProperties(prefix="server.database")
+class DatabaseConfig:
+ """数据库配置类"""
+ host: str = config_field(default="localhost")
+ port: int = config_field(default=3306)
+ username: str = config_field(default="root")
+ password: str = config_field(default="")
+ database: str = config_field(default="xiaozhi")
+
+
+@ConfigurationProperties(prefix="server.redis")
+class RedisConfig:
+ """Redis配置类"""
+ host: str = config_field(default="localhost")
+ port: int = config_field(default=6379)
+ password: str = config_field(default="")
+ db: int = config_field(default=0)
+
+
+@ConfigurationProperties(prefix="mqtt_server")
+class MQTTServerConfig:
+ """MQTT服务器配置类"""
+ enabled: bool = config_field(default=False)
+ host: str = config_field(default="0.0.0.0")
+ port: int = config_field(default=1883)
+ udp_port: int = config_field(default=1883)
+ public_ip: str = config_field(default="localhost")
+ max_connections: int = config_field(default=1000)
+ heartbeat_interval: int = config_field(default=30)
+ max_payload_size: int = config_field(default=8192)
+
+
+@ConfigurationProperties(prefix="server")
+class ServerConfig:
+ """服务器配置类"""
+ ip: str = config_field(default="0.0.0.0")
+ port: int = config_field(default=8080)
+ http_port: int = config_field(default=8081)
+ auth_key: str = config_field(default="")
+ vision_explain: str = config_field(default="")
+
+ # 嵌套配置类
+ database: DatabaseConfig = config_field(default_factory=lambda: DatabaseConfig())
+ redis: RedisConfig = config_field(default_factory=lambda: RedisConfig())
+ mqtt_server: MQTTServerConfig = config_field(default_factory=lambda: MQTTServerConfig())
+
+
+@ConfigurationProperties(prefix="asr.whisper")
+class WhisperConfig:
+ """Whisper ASR配置类"""
+ model: str = config_field(default="base")
+ language: str = config_field(default="zh")
+ device: str = config_field(default="cpu")
+
+
+@ConfigurationProperties(prefix="asr")
+class ASRConfig:
+ """ASR配置类"""
+ provider: str = config_field(default="whisper")
+ # 嵌套配置
+ whisper: WhisperConfig = config_field(default_factory=lambda: WhisperConfig())
+
+
+@ConfigurationProperties(prefix="selected_module")
+class SelectedModuleConfig:
+ """选中模块配置类"""
+ ASR: str = config_field(default="")
+ TTS: str = config_field(default="")
+ LLM: str = config_field(default="")
+ VLLM: str = config_field(default="")
+ VAD: str = config_field(default="")
+ Memory: str = config_field(default="")
+ Intent: str = config_field(default="")
+
+
+@ConfigurationProperties(prefix="log")
+class LogConfig:
+ """日志配置类"""
+ log_dir: str = config_field(default="tmp")
+ level: str = config_field(default="INFO")
+
+
+@ConfigurationProperties(prefix="protocols")
+class ProtocolConfig:
+ """协议配置类"""
+ enabled_protocols: list = config_field(default_factory=lambda: ["websocket"])
+ websocket_enabled: bool = config_field(default=True)
+ mqtt_enabled: bool = config_field(default=False)
+
+
+@ConfigurationProperties(prefix="")
+class MainConfig:
+ """主配置类,包含常用的顶级配置"""
+ read_config_from_api: bool = config_field(default=False)
+ exit_commands: list = config_field(default_factory=list)
+ close_connection_no_voice_time: int = config_field(default=120)
+ xiaozhi: str = config_field(default="")
+ prompt: str = config_field(default="")
+ delete_audio: bool = config_field(default=True)
+
+ # 协议配置
+ protocols: ProtocolConfig = config_field(default_factory=lambda: ProtocolConfig())
+
+
+@ConfigurationProperties(prefix="voiceprint")
+class VoiceprintConfig:
+ """声纹配置类"""
+ enabled: bool = config_field(default=False)
+ model_path: str = config_field(default="")
+ threshold: float = config_field(default=0.5)
+
+
+# 全局配置实例
+_global_config_dict: ConfigDict = None
+_server_config: ServerConfig = None
+_selected_module_config: SelectedModuleConfig = None
+_log_config: LogConfig = None
+_main_config: MainConfig = None
+_voiceprint_config: VoiceprintConfig = None
+_mqtt_server_config: MQTTServerConfig = None
+_protocol_config: ProtocolConfig = None
+
+
+def get_config_instance(config_class: Type[T]) -> T:
+ """获取配置类实例的工厂方法"""
+ global _global_config_dict
+ if _global_config_dict is None:
+ _global_config_dict = load_config()
+
+ return config_class(_global_config_dict)
+
+
+def get_server_config() -> ServerConfig:
+ """获取服务器配置实例"""
+ global _server_config
+ if _server_config is None:
+ _server_config = get_config_instance(ServerConfig)
+ return _server_config
+
+
+def get_selected_module_config() -> SelectedModuleConfig:
+ """获取选中模块配置实例"""
+ global _selected_module_config
+ if _selected_module_config is None:
+ _selected_module_config = get_config_instance(SelectedModuleConfig)
+ return _selected_module_config
+
+
+def get_log_config() -> LogConfig:
+ """获取日志配置实例"""
+ global _log_config
+ if _log_config is None:
+ _log_config = get_config_instance(LogConfig)
+ return _log_config
+
+
+def get_main_config() -> MainConfig:
+ """获取主配置实例"""
+ global _main_config
+ if _main_config is None:
+ _main_config = get_config_instance(MainConfig)
+ return _main_config
+
+
+def get_voiceprint_config() -> VoiceprintConfig:
+ """获取声纹配置实例"""
+ global _voiceprint_config
+ if _voiceprint_config is None:
+ _voiceprint_config = get_config_instance(VoiceprintConfig)
+ return _voiceprint_config
+
+
+def get_mqtt_server_config() -> MQTTServerConfig:
+ """获取MQTT服务器配置实例"""
+ global _mqtt_server_config
+ if _mqtt_server_config is None:
+ _mqtt_server_config = get_config_instance(MQTTServerConfig)
+ return _mqtt_server_config
+
+
+def get_protocol_config() -> ProtocolConfig:
+ """获取协议配置实例"""
+ global _protocol_config
+ if _protocol_config is None:
+ _protocol_config = get_config_instance(ProtocolConfig)
+ return _protocol_config
+
+
+def refresh_config():
+ """刷新所有配置实例"""
+ global _global_config_dict, _server_config, _selected_module_config, _log_config, _main_config, _voiceprint_config, _mqtt_server_config, _protocol_config
+ _global_config_dict = None
+ _server_config = None
+ _selected_module_config = None
+ _log_config = None
+ _main_config = None
+ _voiceprint_config = None
+ _mqtt_server_config = None
+ _protocol_config = None
diff --git a/main/xiaozhi-server/core/components/adapters/asr_adapter.py b/main/xiaozhi-server/core/components/adapters/asr_adapter.py
new file mode 100644
index 00000000..776e81b6
--- /dev/null
+++ b/main/xiaozhi-server/core/components/adapters/asr_adapter.py
@@ -0,0 +1,76 @@
+from typing import Any, Dict
+from core.components.component_manager import Component, ComponentType, ComponentFactory
+from core.utils import asr
+from core.utils.modules_initialize import initialize_asr
+from config.logger import setup_logging
+
+logger = setup_logging()
+
+
+class ASRAdapter(Component):
+ """ASR组件适配器:将现有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)
+
+ async def _do_initialize(self, context: Any) -> None:
+ """初始化ASR组件"""
+ try:
+ # 获取ASR配置
+ selected_module = self.config.get("selected_module", {}).get("ASR")
+ if not selected_module:
+ raise ValueError("未配置ASR模块")
+
+ # 创建ASR实例
+ self._asr_instance = initialize_asr(self.config)
+
+ # 注册资源以便清理
+ 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}")
+
+ except Exception as e:
+ logger.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组件清理完成")
+
+ except Exception as e:
+ logger.error(f"ASR组件清理失败: {e}")
+ finally:
+ self._asr_instance = None
+
+ @property
+ def asr_instance(self):
+ """获取ASR实例"""
+ return self._asr_instance
+
+
+class ASRFactory(ComponentFactory):
+ """ASR组件工厂"""
+
+ def create(self, config: Dict[str, Any]) -> Component:
+ return ASRAdapter(config)
+
+ def get_component_type(self) -> ComponentType:
+ return ComponentType.ASR
+
+
diff --git a/main/xiaozhi-server/core/components/adapters/intent_adapter.py b/main/xiaozhi-server/core/components/adapters/intent_adapter.py
new file mode 100644
index 00000000..f7fffbac
--- /dev/null
+++ b/main/xiaozhi-server/core/components/adapters/intent_adapter.py
@@ -0,0 +1,85 @@
+from typing import Any, Dict
+from core.components.component_manager import Component, ComponentType, ComponentFactory
+from core.utils import intent
+from config.logger import setup_logging
+
+logger = setup_logging()
+
+
+class IntentAdapter(Component):
+ """Intent组件适配器:将现有Intent组件包装为新的组件接口"""
+
+ def __init__(self, config: Dict[str, Any]):
+ super().__init__(ComponentType.INTENT, config)
+ self._intent_instance = None
+
+ async def _do_initialize(self, context: Any) -> None:
+ """初始化Intent组件"""
+ try:
+ # 获取Intent配置
+ selected_module = self.config.get("selected_module", {}).get("Intent")
+ if not selected_module:
+ raise ValueError("未配置Intent模块")
+
+ # 获取Intent类型
+ intent_type = (
+ selected_module
+ if "type" not in self.config["Intent"][selected_module]
+ else self.config["Intent"][selected_module]["type"]
+ )
+
+ # 创建Intent实例
+ self._intent_instance = intent.create_instance(
+ intent_type,
+ self.config["Intent"][selected_module],
+ )
+
+ # 注册资源以便清理
+ self.add_resource(self._intent_instance)
+
+ # 设置LLM(如果需要)
+ if intent_type in ["intent_llm", "function_call"]:
+ llm_component = context.components.get('llm')
+ if llm_component and hasattr(llm_component, 'llm_instance'):
+ if hasattr(self._intent_instance, 'set_llm'):
+ self._intent_instance.set_llm(llm_component.llm_instance)
+
+ logger.info(f"Intent组件初始化完成: {intent_type}")
+
+ except Exception as e:
+ logger.error(f"Intent组件初始化失败: {e}")
+ raise
+
+ async def _do_cleanup(self) -> None:
+ """清理Intent组件"""
+ if self._intent_instance:
+ try:
+ # 关闭Intent实例
+ if hasattr(self._intent_instance, 'close'):
+ await self._intent_instance.close()
+ elif hasattr(self._intent_instance, 'cleanup'):
+ await self._intent_instance.cleanup()
+
+ logger.info("Intent组件清理完成")
+
+ except Exception as e:
+ logger.error(f"Intent组件清理失败: {e}")
+ finally:
+ self._intent_instance = None
+
+ @property
+ def intent_instance(self):
+ """获取Intent实例"""
+ return self._intent_instance
+
+
+class IntentFactory(ComponentFactory):
+ """Intent组件工厂"""
+
+ def create(self, config: Dict[str, Any]) -> Component:
+ return IntentAdapter(config)
+
+ def get_component_type(self) -> ComponentType:
+ return ComponentType.INTENT
+
+
diff --git a/main/xiaozhi-server/core/components/adapters/llm_adapter.py b/main/xiaozhi-server/core/components/adapters/llm_adapter.py
new file mode 100644
index 00000000..9fd4968e
--- /dev/null
+++ b/main/xiaozhi-server/core/components/adapters/llm_adapter.py
@@ -0,0 +1,78 @@
+from typing import Any, Dict
+from core.components.component_manager import Component, ComponentType, ComponentFactory
+from core.utils import llm
+from config.logger import setup_logging
+
+logger = setup_logging()
+
+
+class LLMAdapter(Component):
+ """LLM组件适配器:将现有LLM组件包装为新的组件接口"""
+
+ def __init__(self, config: Dict[str, Any]):
+ super().__init__(ComponentType.LLM, config)
+ self._llm_instance = None
+
+ async def _do_initialize(self, context: Any) -> None:
+ """初始化LLM组件"""
+ try:
+ # 获取LLM配置
+ selected_module = self.config.get("selected_module", {}).get("LLM")
+ if not selected_module:
+ raise ValueError("未配置LLM模块")
+
+ # 获取LLM类型
+ llm_type = (
+ selected_module
+ if "type" not in self.config["LLM"][selected_module]
+ else self.config["LLM"][selected_module]["type"]
+ )
+
+ # 创建LLM实例
+ self._llm_instance = llm.create_instance(
+ llm_type,
+ self.config["LLM"][selected_module],
+ )
+
+ # 注册资源以便清理
+ self.add_resource(self._llm_instance)
+
+ logger.info(f"LLM组件初始化完成: {llm_type}")
+
+ except Exception as e:
+ logger.error(f"LLM组件初始化失败: {e}")
+ raise
+
+ async def _do_cleanup(self) -> None:
+ """清理LLM组件"""
+ if self._llm_instance:
+ try:
+ # 关闭LLM实例
+ if hasattr(self._llm_instance, 'close'):
+ await self._llm_instance.close()
+ elif hasattr(self._llm_instance, 'cleanup'):
+ await self._llm_instance.cleanup()
+
+ logger.info("LLM组件清理完成")
+
+ except Exception as e:
+ logger.error(f"LLM组件清理失败: {e}")
+ finally:
+ self._llm_instance = None
+
+ @property
+ def llm_instance(self):
+ """获取LLM实例"""
+ return self._llm_instance
+
+
+class LLMFactory(ComponentFactory):
+ """LLM组件工厂"""
+
+ def create(self, config: Dict[str, Any]) -> Component:
+ return LLMAdapter(config)
+
+ def get_component_type(self) -> ComponentType:
+ return ComponentType.LLM
+
+
diff --git a/main/xiaozhi-server/core/components/adapters/memory_adapter.py b/main/xiaozhi-server/core/components/adapters/memory_adapter.py
new file mode 100644
index 00000000..90e5a893
--- /dev/null
+++ b/main/xiaozhi-server/core/components/adapters/memory_adapter.py
@@ -0,0 +1,96 @@
+from typing import Any, Dict
+from core.components.component_manager import Component, ComponentType, ComponentFactory
+from core.utils import memory
+from config.logger import setup_logging
+
+logger = setup_logging()
+
+
+class MemoryAdapter(Component):
+ """Memory组件适配器:将现有Memory组件包装为新的组件接口"""
+
+ def __init__(self, config: Dict[str, Any]):
+ super().__init__(ComponentType.MEMORY, config)
+ self._memory_instance = None
+
+ async def _do_initialize(self, context: Any) -> None:
+ """初始化Memory组件"""
+ try:
+ # 获取Memory配置
+ selected_module = self.config.get("selected_module", {}).get("Memory")
+ if not selected_module:
+ raise ValueError("未配置Memory模块")
+
+ # 获取Memory类型
+ memory_type = (
+ selected_module
+ if "type" not in self.config["Memory"][selected_module]
+ else self.config["Memory"][selected_module]["type"]
+ )
+
+ # 创建Memory实例
+ self._memory_instance = memory.create_instance(
+ memory_type,
+ self.config["Memory"][selected_module],
+ self.config.get("summaryMemory", None),
+ )
+
+ # 注册资源以便清理
+ self.add_resource(self._memory_instance)
+
+ # 初始化记忆模块
+ if hasattr(self._memory_instance, 'init_memory'):
+ # 需要LLM实例来初始化记忆
+ llm_component = context.components.get('llm')
+ if llm_component and hasattr(llm_component, 'llm_instance'):
+ self._memory_instance.init_memory(
+ role_id=context.device_id,
+ llm=llm_component.llm_instance,
+ summary_memory=self.config.get("summaryMemory", None),
+ save_to_file=not self.config.get("read_config_from_api", False),
+ )
+
+ logger.info(f"Memory组件初始化完成: {memory_type}")
+
+ except Exception as e:
+ logger.error(f"Memory组件初始化失败: {e}")
+ raise
+
+ async def _do_cleanup(self) -> None:
+ """清理Memory组件"""
+ if self._memory_instance:
+ try:
+ # 保存记忆
+ if hasattr(self._memory_instance, 'save_memory'):
+ # 这里需要获取对话历史,暂时跳过
+ pass
+
+ # 关闭Memory实例
+ if hasattr(self._memory_instance, 'close'):
+ await self._memory_instance.close()
+ elif hasattr(self._memory_instance, 'cleanup'):
+ await self._memory_instance.cleanup()
+
+ logger.info("Memory组件清理完成")
+
+ except Exception as e:
+ logger.error(f"Memory组件清理失败: {e}")
+ finally:
+ self._memory_instance = None
+
+ @property
+ def memory_instance(self):
+ """获取Memory实例"""
+ return self._memory_instance
+
+
+class MemoryFactory(ComponentFactory):
+ """Memory组件工厂"""
+
+ def create(self, config: Dict[str, Any]) -> Component:
+ return MemoryAdapter(config)
+
+ def get_component_type(self) -> ComponentType:
+ return ComponentType.MEMORY
+
+
diff --git a/main/xiaozhi-server/core/components/adapters/tts_adapter.py b/main/xiaozhi-server/core/components/adapters/tts_adapter.py
new file mode 100644
index 00000000..23e80db3
--- /dev/null
+++ b/main/xiaozhi-server/core/components/adapters/tts_adapter.py
@@ -0,0 +1,90 @@
+from typing import Any, Dict
+from core.components.component_manager import Component, ComponentType, ComponentFactory
+from core.utils import tts
+from config.logger import setup_logging
+
+logger = setup_logging()
+
+
+class TTSAdapter(Component):
+ """TTS组件适配器:将现有TTS组件包装为新的组件接口"""
+
+ def __init__(self, config: Dict[str, Any]):
+ super().__init__(ComponentType.TTS, config)
+ self._tts_instance = None
+ self._delete_audio = config.get("delete_audio", True)
+
+ async def _do_initialize(self, context: Any) -> None:
+ """初始化TTS组件"""
+ try:
+ # 获取TTS配置
+ selected_module = self.config.get("selected_module", {}).get("TTS")
+ if not selected_module:
+ raise ValueError("未配置TTS模块")
+
+ # 获取TTS类型
+ tts_type = (
+ selected_module
+ if "type" not in self.config["TTS"][selected_module]
+ else self.config["TTS"][selected_module]["type"]
+ )
+
+ # 创建TTS实例
+ self._tts_instance = tts.create_instance(
+ tts_type,
+ self.config["TTS"][selected_module],
+ str(self._delete_audio).lower() in ("true", "1", "yes"),
+ )
+
+ # 注册资源以便清理
+ self.add_resource(self._tts_instance)
+
+ # 打开音频通道
+ if hasattr(self._tts_instance, 'open_audio_channels'):
+ await self._tts_instance.open_audio_channels(context)
+
+ # 设置兼容属性(用于向后兼容)
+ if hasattr(context, 'tts'):
+ context.tts = self._tts_instance
+
+ logger.info(f"TTS组件初始化完成: {tts_type}")
+
+ except Exception as e:
+ logger.error(f"TTS组件初始化失败: {e}")
+ raise
+
+ async def _do_cleanup(self) -> None:
+ """清理TTS组件"""
+ if self._tts_instance:
+ try:
+ # 关闭TTS实例
+ if hasattr(self._tts_instance, 'close'):
+ await self._tts_instance.close()
+
+ # 清理音频文件
+ if hasattr(self._tts_instance, 'cleanup_audio_files'):
+ self._tts_instance.cleanup_audio_files()
+
+ logger.info("TTS组件清理完成")
+
+ except Exception as e:
+ logger.error(f"TTS组件清理失败: {e}")
+ finally:
+ self._tts_instance = None
+
+ @property
+ def tts_instance(self):
+ """获取TTS实例"""
+ return self._tts_instance
+
+
+class TTSFactory(ComponentFactory):
+ """TTS组件工厂"""
+
+ def create(self, config: Dict[str, Any]) -> Component:
+ return TTSAdapter(config)
+
+ def get_component_type(self) -> ComponentType:
+ return ComponentType.TTS
+
+
diff --git a/main/xiaozhi-server/core/components/adapters/vad_adapter.py b/main/xiaozhi-server/core/components/adapters/vad_adapter.py
new file mode 100644
index 00000000..738ec0b2
--- /dev/null
+++ b/main/xiaozhi-server/core/components/adapters/vad_adapter.py
@@ -0,0 +1,78 @@
+from typing import Any, Dict
+from core.components.component_manager import Component, ComponentType, ComponentFactory
+from core.utils import vad
+from config.logger import setup_logging
+
+logger = setup_logging()
+
+
+class VADAdapter(Component):
+ """VAD组件适配器:将现有VAD组件包装为新的组件接口"""
+
+ def __init__(self, config: Dict[str, Any]):
+ super().__init__(ComponentType.VAD, config)
+ self._vad_instance = None
+
+ async def _do_initialize(self, context: Any) -> None:
+ """初始化VAD组件"""
+ try:
+ # 获取VAD配置
+ selected_module = self.config.get("selected_module", {}).get("VAD")
+ if not selected_module:
+ raise ValueError("未配置VAD模块")
+
+ # 获取VAD类型
+ vad_type = (
+ selected_module
+ if "type" not in self.config["VAD"][selected_module]
+ else self.config["VAD"][selected_module]["type"]
+ )
+
+ # 创建VAD实例
+ self._vad_instance = vad.create_instance(
+ vad_type,
+ self.config["VAD"][selected_module],
+ )
+
+ # 注册资源以便清理
+ self.add_resource(self._vad_instance)
+
+ logger.info(f"VAD组件初始化完成: {vad_type}")
+
+ except Exception as e:
+ logger.error(f"VAD组件初始化失败: {e}")
+ raise
+
+ async def _do_cleanup(self) -> None:
+ """清理VAD组件"""
+ if self._vad_instance:
+ try:
+ # 关闭VAD实例
+ if hasattr(self._vad_instance, 'close'):
+ await self._vad_instance.close()
+ elif hasattr(self._vad_instance, 'cleanup'):
+ await self._vad_instance.cleanup()
+
+ logger.info("VAD组件清理完成")
+
+ except Exception as e:
+ logger.error(f"VAD组件清理失败: {e}")
+ finally:
+ self._vad_instance = None
+
+ @property
+ def vad_instance(self):
+ """获取VAD实例"""
+ return self._vad_instance
+
+
+class VADFactory(ComponentFactory):
+ """VAD组件工厂"""
+
+ def create(self, config: Dict[str, Any]) -> Component:
+ return VADAdapter(config)
+
+ def get_component_type(self) -> ComponentType:
+ return ComponentType.VAD
+
+
diff --git a/main/xiaozhi-server/core/components/component_manager.py b/main/xiaozhi-server/core/components/component_manager.py
new file mode 100644
index 00000000..009fdd5d
--- /dev/null
+++ b/main/xiaozhi-server/core/components/component_manager.py
@@ -0,0 +1,235 @@
+import asyncio
+import weakref
+from abc import ABC, abstractmethod
+from typing import Any, Callable, Dict, Optional, Type, TypeVar, Generic
+from enum import Enum
+from config.logger import setup_logging
+
+T = TypeVar('T')
+
+logger = setup_logging()
+
+
+class ComponentType(Enum):
+ """组件类型枚举"""
+ TTS = "tts"
+ ASR = "asr"
+ VAD = "vad"
+ LLM = "llm"
+ MEMORY = "memory"
+ INTENT = "intent"
+
+
+class ComponentState(Enum):
+ """组件状态枚举"""
+ UNINITIALIZED = "uninitialized"
+ INITIALIZING = "initializing"
+ READY = "ready"
+ ERROR = "error"
+ CLEANING = "cleaning"
+ CLEANED = "cleaned"
+
+
+class Component(ABC):
+ """组件基类:定义统一的组件接口和生命周期管理"""
+
+ def __init__(self, component_type: ComponentType, config: Dict[str, Any]):
+ self.component_type = component_type
+ self.config = config
+ self.state = ComponentState.UNINITIALIZED
+ self._initialization_lock = asyncio.Lock()
+ self._cleanup_lock = asyncio.Lock()
+ self._dependencies: Dict[str, 'Component'] = {}
+ self._dependents: weakref.WeakSet['Component'] = weakref.WeakSet()
+ self._resources: list = [] # 存储需要清理的资源
+
+ @abstractmethod
+ async def _do_initialize(self, context: Any) -> None:
+ """子类实现具体的初始化逻辑"""
+ pass
+
+ @abstractmethod
+ async def _do_cleanup(self) -> None:
+ """子类实现具体的清理逻辑"""
+ pass
+
+ async def initialize(self, context: Any) -> None:
+ """初始化组件(带锁保护)"""
+ async with self._initialization_lock:
+ if self.state != ComponentState.UNINITIALIZED:
+ return
+
+ try:
+ self.state = ComponentState.INITIALIZING
+ logger.info(f"正在初始化组件: {self.component_type.value}")
+
+ # 初始化依赖组件
+ await self._initialize_dependencies(context)
+
+ # 执行具体初始化
+ await self._do_initialize(context)
+
+ self.state = ComponentState.READY
+ logger.info(f"组件初始化完成: {self.component_type.value}")
+
+ except Exception as e:
+ self.state = ComponentState.ERROR
+ logger.error(f"组件初始化失败: {self.component_type.value}, 错误: {e}")
+ raise
+
+ async def cleanup(self) -> None:
+ """清理组件(带锁保护)"""
+ async with self._cleanup_lock:
+ if self.state in [ComponentState.CLEANING, ComponentState.CLEANED]:
+ return
+
+ try:
+ self.state = ComponentState.CLEANING
+ logger.info(f"正在清理组件: {self.component_type.value}")
+
+ # 清理依赖此组件的其他组件
+ await self._cleanup_dependents()
+
+ # 执行具体清理
+ await self._do_cleanup()
+
+ # 清理资源
+ await self._cleanup_resources()
+
+ self.state = ComponentState.CLEANED
+ logger.info(f"组件清理完成: {self.component_type.value}")
+
+ except Exception as e:
+ logger.error(f"组件清理失败: {self.component_type.value}, 错误: {e}")
+ # 即使清理失败,也标记为已清理,避免重复清理
+ self.state = ComponentState.CLEANED
+
+ def add_dependency(self, name: str, component: 'Component') -> None:
+ """添加依赖组件"""
+ self._dependencies[name] = component
+ component._dependents.add(self)
+
+ def add_resource(self, resource: Any) -> None:
+ """添加需要清理的资源"""
+ self._resources.append(resource)
+
+ async def _initialize_dependencies(self, context: Any) -> None:
+ """初始化依赖组件"""
+ for name, dep in self._dependencies.items():
+ if dep.state == ComponentState.UNINITIALIZED:
+ await dep.initialize(context)
+
+ async def _cleanup_dependents(self) -> None:
+ """清理依赖此组件的其他组件"""
+ for dependent in list(self._dependents):
+ await dependent.cleanup()
+
+ async def _cleanup_resources(self) -> None:
+ """清理所有注册的资源"""
+ for resource in self._resources:
+ try:
+ if hasattr(resource, 'close'):
+ if asyncio.iscoroutinefunction(resource.close):
+ await resource.close()
+ else:
+ resource.close()
+ elif hasattr(resource, 'cleanup'):
+ if asyncio.iscoroutinefunction(resource.cleanup):
+ await resource.cleanup()
+ else:
+ resource.cleanup()
+ except Exception as e:
+ logger.warning(f"清理资源时出错: {e}")
+ self._resources.clear()
+
+
+class ComponentFactory(ABC):
+ """组件工厂基类"""
+
+ @abstractmethod
+ def create(self, config: Dict[str, Any]) -> Component:
+ """创建组件实例"""
+ pass
+
+ @abstractmethod
+ def get_component_type(self) -> ComponentType:
+ """获取组件类型"""
+ pass
+
+
+class ComponentManager:
+ """
+ 组件管理器:统一管理连接期内的组件实例生命周期。
+ 支持分类管理、依赖注入、按需懒加载与统一清理。
+ """
+
+ def __init__(self, config: Dict[str, Any]):
+ self._config = config
+ self._components: Dict[str, Component] = {}
+ self._factories: Dict[ComponentType, ComponentFactory] = {}
+ self._initialization_order: list[ComponentType] = []
+
+ def register_factory(self, factory: ComponentFactory) -> None:
+ """注册组件工厂"""
+ component_type = factory.get_component_type()
+ self._factories[component_type] = factory
+ logger.debug(f"注册组件工厂: {component_type.value}")
+
+ def set_initialization_order(self, order: list[ComponentType]) -> None:
+ """设置组件初始化顺序"""
+ self._initialization_order = order
+
+ async def get_component(self, component_type: ComponentType, context: Any) -> Optional[Component]:
+ """获取组件实例(按需创建)"""
+ key = component_type.value
+
+ if key not in self._components:
+ factory = self._factories.get(component_type)
+ if factory is None:
+ logger.warning(f"未找到组件工厂: {component_type.value}")
+ return None
+
+ try:
+ instance = factory.create(self._config)
+ await instance.initialize(context)
+ self._components[key] = instance
+ logger.info(f"组件创建并初始化完成: {component_type.value}")
+ except Exception as e:
+ logger.error(f"组件创建失败: {component_type.value}, 错误: {e}")
+ return None
+
+ return self._components.get(key)
+
+ def get(self, component_name: str) -> Optional[Component]:
+ """获取已初始化的组件实例(兼容接口)"""
+ return self._components.get(component_name)
+
+ async def initialize_all(self, context: Any) -> None:
+ """按顺序初始化所有组件"""
+ for component_type in self._initialization_order:
+ await self.get_component(component_type, context)
+
+ async def cleanup_all(self) -> None:
+ """清理所有组件(逆序清理)"""
+ # 按逆序清理,确保依赖关系正确
+ for component_type in reversed(self._initialization_order):
+ key = component_type.value
+ if key in self._components:
+ component = self._components[key]
+ await component.cleanup()
+ del self._components[key]
+
+ # 清理可能遗漏的组件
+ remaining_components = list(self._components.values())
+ for component in remaining_components:
+ await component.cleanup()
+
+ self._components.clear()
+ logger.info("所有组件已清理完成")
+
+ def get_component_status(self) -> Dict[str, str]:
+ """获取所有组件状态"""
+ return {
+ name: component.state.value
+ for name, component in self._components.items()
+ }
\ No newline at end of file
diff --git a/main/xiaozhi-server/core/components/component_registry.py b/main/xiaozhi-server/core/components/component_registry.py
new file mode 100644
index 00000000..9bcd68f3
--- /dev/null
+++ b/main/xiaozhi-server/core/components/component_registry.py
@@ -0,0 +1,75 @@
+from typing import Dict, Any
+from core.components.component_manager import ComponentManager, ComponentType
+from core.components.adapters.tts_adapter import TTSFactory
+from core.components.adapters.asr_adapter import ASRFactory
+from core.components.adapters.vad_adapter import VADFactory
+from core.components.adapters.llm_adapter import LLMFactory
+from core.components.adapters.memory_adapter import MemoryFactory
+from core.components.adapters.intent_adapter import IntentFactory
+from config.logger import setup_logging
+
+logger = setup_logging()
+
+
+class ComponentRegistry:
+ """组件注册器:统一管理所有组件工厂的注册"""
+
+ _factories_registered = False
+
+ @classmethod
+ def create_component_manager(cls, config: Dict[str, Any]) -> ComponentManager:
+ """创建并配置组件管理器"""
+ manager = ComponentManager(config)
+
+ # 只在第一次时记录注册日志
+ if not cls._factories_registered:
+ logger.info("注册组件工厂")
+ cls._factories_registered = True
+
+ # 注册所有组件工厂(每个manager都需要注册,但不重复记录日志)
+ manager.register_factory(TTSFactory())
+ manager.register_factory(ASRFactory())
+ manager.register_factory(VADFactory())
+ manager.register_factory(LLMFactory())
+ manager.register_factory(MemoryFactory())
+ manager.register_factory(IntentFactory())
+
+ # 设置组件初始化顺序(考虑依赖关系)
+ # VAD -> ASR -> LLM -> Memory/Intent -> TTS
+ initialization_order = [
+ ComponentType.VAD,
+ ComponentType.ASR,
+ ComponentType.LLM,
+ ComponentType.MEMORY,
+ ComponentType.INTENT,
+ ComponentType.TTS,
+ ]
+ manager.set_initialization_order(initialization_order)
+
+ if not cls._factories_registered:
+ logger.info("组件管理器创建完成,已注册所有组件工厂")
+
+ return manager
+
+ @staticmethod
+ def get_required_components(config: Dict[str, Any]) -> list[ComponentType]:
+ """根据配置获取需要的组件类型"""
+ required = []
+ selected_modules = config.get("selected_module", {})
+
+ if selected_modules.get("VAD"):
+ required.append(ComponentType.VAD)
+ if selected_modules.get("ASR"):
+ required.append(ComponentType.ASR)
+ if selected_modules.get("LLM"):
+ required.append(ComponentType.LLM)
+ if selected_modules.get("TTS"):
+ required.append(ComponentType.TTS)
+ if selected_modules.get("Memory"):
+ required.append(ComponentType.MEMORY)
+ if selected_modules.get("Intent"):
+ required.append(ComponentType.INTENT)
+
+ return required
+
+
diff --git a/main/xiaozhi-server/core/connection.py b/main/xiaozhi-server/core/connection.py
index 4b3b9d1c..7e94c9cd 100644
--- a/main/xiaozhi-server/core/connection.py
+++ b/main/xiaozhi-server/core/connection.py
@@ -23,23 +23,25 @@ from core.utils.modules_initialize import (
initialize_tts,
initialize_asr,
)
-from core.handle.reportHandle import report
+# from core.handle.reportHandle import report # 旧架构,已被新架构替代
from core.providers.tts.default import DefaultTTS
from concurrent.futures import ThreadPoolExecutor
from core.utils.dialogue import Message, Dialogue
from core.providers.asr.dto.dto import InterfaceType
-from core.handle.textHandle import handleTextMessage
+# from core.handle.textHandle import handleTextMessage # 旧架构,已被新架构替代
from core.providers.tools.unified_tool_handler import UnifiedToolHandler
from plugins_func.loadplugins import auto_import_modules
from plugins_func.register import Action, ActionResponse
from core.auth import AuthMiddleware, AuthenticationError
-from config.config_loader import get_private_config_from_api
+from config.config_loader import get_private_config_from_api, ConfigDict
from core.providers.tts.dto.dto import ContentType, TTSMessageDTO, SentenceType
from config.logger import setup_logging, build_module_string, create_connection_logger
from config.manage_api_client import DeviceNotFoundException, DeviceBindException
from core.utils.prompt_manager import PromptManager
from core.utils.voiceprint_provider import VoiceprintProvider
from core.utils import textUtils
+from core.context.session_context import SessionContext
+from core.components.component_registry import ComponentRegistry
TAG = __name__
@@ -139,7 +141,7 @@ class ConnectionHandler:
self.iot_descriptors = {}
self.func_handler = None
- self.cmd_exit = self.config["exit_commands"]
+ self.cmd_exit = self.config.get("exit_commands", [])
# 是否在聊天结束后关闭连接
self.close_after_chat = False
@@ -156,6 +158,11 @@ class ConnectionHandler:
# 初始化提示词管理器
self.prompt_manager = PromptManager(config, self.logger)
+
+ # 新增:会话上下文与组件管理器(会话级清理)
+ self.session_context: SessionContext = SessionContext()
+ self.session_context.config = self.config
+ self.component_manager = ComponentRegistry.create_component_manager(self.config)
async def handle_connection(self, ws):
try:
@@ -197,6 +204,10 @@ class ConnectionHandler:
# 认证通过,继续处理
self.websocket = ws
self.device_id = self.headers.get("device-id", None)
+ # 更新会话上下文关键信息
+ self.session_context.headers = self.headers
+ self.session_context.device_id = self.device_id
+ self.session_context.client_ip = self.client_ip
# 初始化活动时间戳
self.last_activity_time = time.time() * 1000
@@ -204,7 +215,8 @@ class ConnectionHandler:
# 启动超时检查任务
self.timeout_task = asyncio.create_task(self._check_timeout())
- self.welcome_msg = self.config["xiaozhi"]
+ # 新的配置访问方式 - 使用点号访问
+ self.welcome_msg = self.config.xiaozhi
self.welcome_msg["session_id"] = self.session_id
# 获取差异化配置
@@ -237,6 +249,18 @@ class ConnectionHandler:
self.logger.bind(tag=TAG).error(
f"强制关闭连接时出错: {close_error}"
)
+ finally:
+ # 会话级组件与回调清理(容错)
+ try:
+ if hasattr(self, "component_manager") and self.component_manager:
+ await self.component_manager.cleanup_all()
+ except Exception as e:
+ self.logger.bind(tag=TAG).error(f"组件清理失败: {e}")
+ try:
+ if hasattr(self, "session_context") and self.session_context:
+ await self.session_context.run_cleanup()
+ except Exception as e:
+ self.logger.bind(tag=TAG).error(f"会话清理回调执行失败: {e}")
async def _save_and_close(self, ws):
"""保存记忆并关闭连接"""
@@ -479,7 +503,8 @@ class ConnectionHandler:
if init_vad:
self.config["VAD"] = private_config["VAD"]
- self.config["selected_module"]["VAD"] = private_config["selected_module"][
+ # 新的配置访问方式 - 使用嵌套路径设置
+ self.config["selected_module.VAD"] = private_config["selected_module"][
"VAD"
]
if init_asr:
@@ -577,17 +602,16 @@ class ConnectionHandler:
# 获取记忆总结配置
memory_config = self.config["Memory"]
- memory_type = self.config["Memory"][self.config["selected_module"]["Memory"]][
- "type"
- ]
+ # 新的配置访问方式 - 使用嵌套get方法
+ memory_module = self.config.get("selected_module.Memory")
+ memory_type = self.config.get(f"Memory.{memory_module}.type")
# 如果使用 nomen,直接返回
if memory_type == "nomem":
return
# 使用 mem_local_short 模式
elif memory_type == "mem_local_short":
- memory_llm_name = memory_config[self.config["selected_module"]["Memory"]][
- "llm"
- ]
+ # 新的配置访问方式 - 使用嵌套get和f-string
+ memory_llm_name = self.config.get(f"Memory.{memory_module}.llm")
if memory_llm_name and memory_llm_name in self.config["LLM"]:
# 如果配置了专用LLM,则创建独立的LLM实例
from core.utils import llm as llm_utils
@@ -907,7 +931,8 @@ class ConnectionHandler:
"""处理上报任务"""
try:
# 执行上报(传入二进制数据)
- report(self, type, text, audio_data, report_time)
+ # report(self, type, text, audio_data, report_time) # 旧架构,已被新架构替代
+ pass # 新架构中由ReportProcessor处理
except Exception as e:
self.logger.bind(tag=TAG).error(f"上报处理异常: {e}")
finally:
diff --git a/main/xiaozhi-server/core/context/session_context.py b/main/xiaozhi-server/core/context/session_context.py
new file mode 100644
index 00000000..b782190c
--- /dev/null
+++ b/main/xiaozhi-server/core/context/session_context.py
@@ -0,0 +1,378 @@
+import copy
+import uuid
+import time
+import queue
+import asyncio
+import threading
+from dataclasses import dataclass, field
+from typing import Any, Dict, Optional, List, Callable, Awaitable, Union
+from collections import deque
+from concurrent.futures import ThreadPoolExecutor
+from core.utils.dialogue import Dialogue
+from core.auth import AuthMiddleware
+from core.utils.prompt_manager import PromptManager
+from core.utils.voiceprint_provider import VoiceprintProvider
+from config.logger import setup_logging
+
+
+@dataclass
+class SessionContext:
+ """
+ 会话上下文:完全替换ConnectionHandler的所有功能
+ 承载单连接生命周期内的状态、组件、资源管理
+ 与传输层解耦,支持WebSocket/MQTT/UDP等多协议
+ """
+
+ # === 基础标识 ===
+ session_id: str = field(default_factory=lambda: str(uuid.uuid4()))
+ device_id: Optional[str] = None
+ client_ip: Optional[str] = None
+ headers: Dict[str, str] = field(default_factory=dict)
+
+ # === 配置管理 ===
+ config: Dict[str, Any] = field(default_factory=dict)
+ common_config: Dict[str, Any] = field(default_factory=dict)
+ private_config: Dict[str, Any] = field(default_factory=dict)
+ selected_module_str: str = ""
+
+ # === 认证与绑定 ===
+ is_authenticated: bool = False
+ need_bind: bool = False
+ bind_code: Optional[str] = None
+ read_config_from_api: bool = False
+ max_output_size: int = 0
+ chat_history_conf: int = 0
+
+ # === 会话状态 ===
+ is_speaking: bool = False
+ listen_mode: str = "auto"
+ abort_requested: bool = False
+ close_after_chat: bool = False
+ just_woken_up: bool = False
+ load_function_plugin: bool = False
+ intent_type: str = "nointent"
+
+ # === 音频相关 ===
+ audio_format: str = "opus"
+ client_have_voice: bool = False
+ client_voice_stop: bool = False
+ client_audio_buffer: bytearray = field(default_factory=bytearray)
+ client_voice_window: deque = field(default_factory=lambda: deque(maxlen=5))
+ last_is_voice: bool = False
+ audio_flow_control: Dict[str, Any] = field(default_factory=dict)
+
+ # === ASR相关 ===
+ asr_audio: List[bytes] = field(default_factory=list)
+ asr_audio_queue: queue.Queue = field(default_factory=queue.Queue)
+ asr_priority_thread: Optional[threading.Thread] = None
+
+ # === LLM相关 ===
+ llm_finish_task: bool = True
+ dialogue: Optional[Dialogue] = None
+ current_speaker: Optional[str] = None
+ sentence_id: Optional[str] = None
+
+ # === TTS相关 ===
+ tts_MessageText: str = ""
+
+ # === IoT相关 ===
+ iot_descriptors: Dict[str, Any] = field(default_factory=dict)
+ func_handler: Optional[Any] = None
+
+ # === 时间管理 ===
+ last_activity_time_ms: float = field(default_factory=lambda: time.time() * 1000)
+ created_at: float = field(default_factory=lambda: time.time())
+ timeout_seconds: int = 180 # 默认超时时间
+ timeout_task: Optional[asyncio.Task] = None
+
+ # === 组件实例 ===
+ # components属性通过@property方法提供,指向component_manager
+
+ # === 其他状态 ===
+ welcome_msg: Optional[Dict[str, Any]] = None
+ prompt: Optional[str] = None
+ features: Optional[Dict[str, Any]] = None
+ mcp_client: Optional[Any] = None
+ cmd_exit: List[str] = field(default_factory=list)
+
+ # === 线程与并发 ===
+ loop: Optional[asyncio.AbstractEventLoop] = None
+ stop_event: Optional[threading.Event] = None
+ executor: Optional[ThreadPoolExecutor] = None
+
+ # === 队列管理 ===
+ report_queue: queue.Queue = field(default_factory=queue.Queue)
+ report_thread: Optional[threading.Thread] = None
+ report_asr_enable: bool = False
+ report_tts_enable: bool = False
+
+ # === 组件管理器 ===
+ component_manager: Optional[Any] = None
+
+ # === 兼容属性(用于向后兼容TTS处理) ===
+ tts: Optional[Any] = None
+ websocket: Optional[Any] = None # 兼容旧TTS组件
+ transport: Optional[Any] = None # 新的transport接口
+
+ # === 工具类 ===
+ auth: Optional[AuthMiddleware] = None
+ prompt_manager: Optional[PromptManager] = None
+ voiceprint_provider: Optional[VoiceprintProvider] = None
+ server: Optional[Any] = None # WebSocket服务器引用
+
+ # === 会话级清理回调 ===
+ _cleanup_callbacks: List[Callable[[], Union[None, Awaitable[None]]]] = field(default_factory=list)
+
+ def __post_init__(self):
+ """初始化后处理"""
+ # 深拷贝配置避免污染
+ if self.config:
+ self.common_config = self.config
+ self.config = copy.deepcopy(self.config)
+
+ # 从配置中读取相关设置
+ self.read_config_from_api = self.config.get("read_config_from_api", False)
+ self.max_output_size = self.config.get("max_output_size", 0)
+ self.chat_history_conf = self.config.get("chat_history_conf", 0)
+ self.cmd_exit = self.config.get("exit_commands", [])
+ self.timeout_seconds = int(self.config.get("close_connection_no_voice_time", 120)) + 60
+
+ # 初始化认证中间件
+ self.auth = AuthMiddleware(self.config)
+
+ # 初始化提示词管理器
+ self.prompt_manager = PromptManager(self.config, setup_logging())
+
+ # 初始化对话管理
+ if not self.dialogue:
+ self.dialogue = Dialogue()
+
+ # 初始化线程相关
+ if not self.loop:
+ try:
+ self.loop = asyncio.get_event_loop()
+ except RuntimeError:
+ self.loop = asyncio.new_event_loop()
+
+ if not self.stop_event:
+ self.stop_event = threading.Event()
+
+ if not self.executor:
+ self.executor = ThreadPoolExecutor(max_workers=5)
+
+ # 初始化上报设置
+ self.report_asr_enable = self.read_config_from_api
+ self.report_tts_enable = self.read_config_from_api
+
+ def update_activity(self) -> None:
+ """刷新最后活跃时间"""
+ self.last_activity_time_ms = time.time() * 1000
+
+ def clearSpeakStatus(self) -> None:
+ """清除服务端讲话状态(兼容方法)"""
+ self.is_speaking = False
+ logger = setup_logging()
+ logger.debug("清除服务端讲话状态")
+
+ def reset_vad_states(self) -> None:
+ """重置VAD状态(兼容方法)"""
+ self.client_audio_buffer = bytearray()
+ self.client_have_voice = False
+ self.client_voice_stop = False
+ logger = setup_logging()
+ logger.debug("VAD states reset.")
+
+ def is_timeout(self, timeout_seconds: int) -> bool:
+ """检查是否超时"""
+ now_ms = time.time() * 1000
+ return (now_ms - self.last_activity_time_ms) > (timeout_seconds * 1000)
+
+ def register_cleanup(self, callback: Callable[[], Union[None, Awaitable[None]]]) -> None:
+ """注册会话结束时需要执行的清理回调"""
+ self._cleanup_callbacks.append(callback)
+
+ async def run_cleanup(self) -> None:
+ """执行所有注册的清理回调"""
+ logger = setup_logging()
+ logger.info(f"Session {self.session_id} 开始执行会话级清理 ({len(self._cleanup_callbacks)} 个回调)")
+
+ # 停止所有线程
+ if self.stop_event:
+ self.stop_event.set()
+
+ # 关闭线程池
+ if self.executor:
+ self.executor.shutdown(wait=False)
+
+ # 取消超时任务
+ if self.timeout_task and not self.timeout_task.done():
+ self.timeout_task.cancel()
+
+ # 执行清理回调
+ for callback in reversed(self._cleanup_callbacks):
+ try:
+ result = callback()
+ if asyncio.iscoroutine(result):
+ await result
+ except Exception as e:
+ logger.error(f"Session {self.session_id} 清理回调执行失败: {e}", exc_info=True)
+
+ self._cleanup_callbacks.clear()
+ logger.info(f"Session {self.session_id} 会话级清理完成")
+
+ # === 兼容旧代码的属性访问 ===
+ @property
+ def client_is_speaking(self) -> bool:
+ """兼容旧代码的属性名"""
+ return self.is_speaking
+
+ @client_is_speaking.setter
+ def client_is_speaking(self, value: bool):
+ self.is_speaking = value
+
+ @property
+ def client_listen_mode(self) -> str:
+ """兼容旧代码的属性名"""
+ return self.listen_mode
+
+ @client_listen_mode.setter
+ def client_listen_mode(self, value: str):
+ self.listen_mode = value
+
+ @property
+ def client_abort(self) -> bool:
+ """兼容旧代码的属性名"""
+ return self.abort_requested
+
+ @client_abort.setter
+ def client_abort(self, value: bool):
+ self.abort_requested = value
+
+ @property
+ def components(self):
+ """组件访问器(兼容属性)"""
+ return self.component_manager
+
+ @components.setter
+ def components(self, value):
+ """组件设置器(兼容属性)- 实际设置到component_manager"""
+ # 如果尝试设置components,我们忽略它或者给出警告
+ # 因为components应该通过component_manager管理
+ logger = setup_logging()
+ logger.warning("尝试直接设置components属性,请使用component_manager")
+
+ @property
+ def last_activity_time(self) -> float:
+ """兼容旧代码:返回毫秒级时间戳"""
+ return self.last_activity_time_ms
+
+ @last_activity_time.setter
+ def last_activity_time(self, value: float):
+ """兼容旧代码:接受毫秒级时间戳"""
+ self.last_activity_time_ms = value
+
+ # === 日志相关 ===
+ @property
+ def logger(self):
+ """获取日志记录器"""
+ return setup_logging()
+
+ # === 工具方法 ===
+ def get_component(self, component_name: str) -> Optional[Any]:
+ """获取组件实例"""
+ return self.components.get(component_name)
+
+ def set_component(self, component_name: str, component_instance: Any) -> None:
+ """设置组件实例"""
+ if self.component_manager:
+ self.component_manager._components[component_name] = component_instance
+
+ def has_component(self, component_name: str) -> bool:
+ """检查是否有指定组件"""
+ return component_name in self.components
+
+ def clear_audio_buffer(self) -> None:
+ """清空音频缓冲区"""
+ self.client_audio_buffer.clear()
+ self.asr_audio.clear()
+
+ # 清空队列
+ try:
+ while not self.asr_audio_queue.empty():
+ self.asr_audio_queue.get_nowait()
+ except queue.Empty:
+ pass
+
+ def reset_voice_state(self) -> None:
+ """重置语音状态"""
+ self.client_have_voice = False
+ self.client_voice_stop = False
+ self.last_is_voice = False
+ self.client_voice_window.clear()
+
+ def initialize_private_config(self) -> None:
+ """初始化差异化配置(从ConnectionHandler迁移)"""
+ from config.config_loader import get_private_config_from_api
+ from config.manage_api_client import DeviceNotFoundException, DeviceBindException
+
+ if not self.read_config_from_api:
+ return
+
+ try:
+ # 获取设备私有配置
+ private_config = get_private_config_from_api(
+ self.config, self.device_id, self.headers.get("client-id")
+ )
+
+ if private_config:
+ self.private_config = private_config
+ # 合并私有配置到主配置
+ self.config.update(private_config)
+
+ except DeviceNotFoundException:
+ self.logger.error(f"设备 {self.device_id} 未找到")
+ self.need_bind = True
+ except DeviceBindException as e:
+ self.logger.error(f"设备绑定异常: {e}")
+ self.need_bind = True
+ self.bind_code = str(e)
+ except Exception as e:
+ self.logger.error(f"获取私有配置失败: {e}")
+
+ async def initialize_components(self) -> None:
+ """异步初始化组件(从ConnectionHandler迁移)"""
+ if not self.component_manager:
+ return
+
+ try:
+ # 初始化各个组件
+ from core.components.component_registry import ComponentType
+
+ # 按依赖顺序初始化组件
+ component_types = [
+ ComponentType.VAD,
+ ComponentType.ASR,
+ ComponentType.LLM,
+ ComponentType.MEMORY,
+ ComponentType.INTENT,
+ ComponentType.TTS
+ ]
+
+ for component_type in component_types:
+ try:
+ component = await self.component_manager.get_component(component_type, self)
+ if component:
+ self.logger.info(f"组件 {component_type} 初始化成功")
+ except Exception as e:
+ self.logger.error(f"组件 {component_type} 初始化失败: {e}")
+
+ except Exception as e:
+ self.logger.error(f"组件初始化失败: {e}")
+
+ def __str__(self) -> str:
+ return f"SessionContext(session_id={self.session_id}, device_id={self.device_id})"
+
+ def __repr__(self) -> str:
+ return self.__str__()
+
+
diff --git a/main/xiaozhi-server/core/pipeline/message_pipeline.py b/main/xiaozhi-server/core/pipeline/message_pipeline.py
new file mode 100644
index 00000000..ce8806d1
--- /dev/null
+++ b/main/xiaozhi-server/core/pipeline/message_pipeline.py
@@ -0,0 +1,28 @@
+from abc import ABC, abstractmethod
+from typing import Any, List
+
+
+class MessageProcessor(ABC):
+ """消息处理器接口。返回 True 表示已处理并中止后续处理。"""
+
+ @abstractmethod
+ async def process(self, context: Any, transport: Any, message: Any) -> bool:
+ raise NotImplementedError
+
+
+class MessagePipeline:
+ """责任链式消息处理管道。"""
+
+ def __init__(self) -> None:
+ self._processors: List[MessageProcessor] = []
+
+ def add_processor(self, processor: MessageProcessor) -> None:
+ self._processors.append(processor)
+
+ async def process_message(self, context: Any, transport: Any, message: Any) -> None:
+ for processor in self._processors:
+ handled = await processor.process(context, transport, message)
+ if handled:
+ return
+
+
diff --git a/main/xiaozhi-server/core/processors/abort_processor.py b/main/xiaozhi-server/core/processors/abort_processor.py
new file mode 100644
index 00000000..8c19817e
--- /dev/null
+++ b/main/xiaozhi-server/core/processors/abort_processor.py
@@ -0,0 +1,102 @@
+import json
+from typing import Any
+from core.pipeline.message_pipeline import MessageProcessor
+from core.context.session_context import SessionContext
+from core.transport.transport_interface import TransportInterface
+from config.logger import setup_logging
+
+logger = setup_logging()
+
+
+class AbortProcessor(MessageProcessor):
+ """中断消息处理器:完整迁移abortMessageHandler.py和abortHandle.py的所有功能"""
+
+ async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
+ """处理abort类型的消息"""
+ if isinstance(message, str):
+ try:
+ msg_json = json.loads(message)
+ if isinstance(msg_json, dict) and msg_json.get("type") == "abort":
+ await self.handle_abort_message(context, transport, msg_json)
+ return True
+ except json.JSONDecodeError:
+ pass
+ return False
+
+ async def handle_abort_message(self, context: SessionContext, transport: TransportInterface, msg_json: dict):
+ """处理中断消息 - 完整迁移自abortHandle.py的handleAbortMessage"""
+ logger.info("Abort message received")
+
+ # 设置成打断状态,会自动打断llm、tts任务 - 完整迁移原逻辑
+ context.abort_requested = True
+
+ # 清理队列 - 完整迁移原逻辑
+ await self._clear_queues(context)
+
+ # 打断客户端说话状态 - 完整迁移原逻辑
+ await transport.send(json.dumps({
+ "type": "tts",
+ "state": "stop",
+ "session_id": context.session_id
+ }))
+
+ # 清理说话状态 - 完整迁移原逻辑
+ self._clear_speak_status(context)
+
+ logger.info("Abort message received-end")
+
+ async def _clear_queues(self, context: SessionContext):
+ """清理所有队列 - 完整迁移原clear_queues逻辑"""
+ try:
+ # 清理TTS音频队列
+ tts_component = context.components.get('tts')
+ if tts_component and hasattr(tts_component, 'tts_instance'):
+ tts_instance = tts_component.tts_instance
+ if hasattr(tts_instance, 'tts_audio_queue'):
+ try:
+ while not tts_instance.tts_audio_queue.empty():
+ tts_instance.tts_audio_queue.get_nowait()
+ except:
+ pass
+
+ # 清理ASR音频队列
+ context.clear_audio_buffer()
+
+ # 清理其他可能的队列
+ if hasattr(context, 'clear_queues'):
+ context.clear_queues()
+
+ except Exception as e:
+ logger.error(f"清理队列时出错: {e}")
+
+ def _clear_speak_status(self, context: SessionContext):
+ """清理说话状态 - 完整迁移原clearSpeakStatus逻辑"""
+ try:
+ # 清理说话状态
+ context.is_speaking = False
+
+ # 如果有其他说话状态相关的属性,也一并清理
+ if hasattr(context, 'clearSpeakStatus'):
+ context.clearSpeakStatus()
+
+ # 重置相关状态
+ context.client_have_voice = False
+ context.client_voice_stop = True
+
+ except Exception as e:
+ logger.error(f"清理说话状态时出错: {e}")
+
+ async def _send_abort_confirmation(self, transport: TransportInterface, session_id: str):
+ """发送中断确认响应(可选)"""
+ response = {
+ "type": "abort",
+ "status": "success",
+ "message": "中断操作已完成",
+ "session_id": session_id
+ }
+
+ try:
+ await transport.send(json.dumps(response))
+ except Exception as e:
+ logger.error(f"发送中断确认响应失败: {e}")
+
diff --git a/main/xiaozhi-server/core/processors/audio_receive_processor.py b/main/xiaozhi-server/core/processors/audio_receive_processor.py
new file mode 100644
index 00000000..c5d4c46e
--- /dev/null
+++ b/main/xiaozhi-server/core/processors/audio_receive_processor.py
@@ -0,0 +1,277 @@
+import time
+import json
+import asyncio
+from typing import Any
+from core.pipeline.message_pipeline import MessageProcessor
+from core.context.session_context import SessionContext
+from core.transport.transport_interface import TransportInterface
+from core.utils.util import audio_to_data
+from core.utils.output_counter import check_device_output_limit
+from config.logger import setup_logging
+
+logger = setup_logging()
+
+
+class AudioReceiveProcessor(MessageProcessor):
+ """音频接收处理器:完整迁移receiveAudioHandle.py的所有功能"""
+
+ async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
+ """处理音频消息"""
+ if isinstance(message, bytes):
+ await self.handle_audio_message(context, transport, message)
+ return True
+ return False
+
+ async def handle_audio_message(self, context: SessionContext, transport: TransportInterface, audio: bytes):
+ """处理音频消息 - 完整迁移自handleAudioMessage"""
+ # 获取VAD组件
+ vad_component = context.components.get('vad')
+ if not vad_component or not hasattr(vad_component, 'vad_instance'):
+ logger.warning("VAD组件未初始化")
+ return
+
+ vad_instance = vad_component.vad_instance
+
+ # 当前片段是否有人说话
+ have_voice = vad_instance.is_vad(context, audio)
+
+ # 如果设备刚刚被唤醒,短暂忽略VAD检测
+ if have_voice and context.just_woken_up:
+ have_voice = False
+ # 设置一个短暂延迟后恢复VAD检测
+ context.asr_audio.clear()
+ if not hasattr(context, "vad_resume_task") or context.vad_resume_task.done():
+ context.vad_resume_task = asyncio.create_task(self._resume_vad_detection(context))
+ return
+
+ if have_voice:
+ if context.is_speaking:
+ await self._handle_abort_message(context, transport)
+
+ # 设备长时间空闲检测,用于say goodbye
+ await self._no_voice_close_connect(context, transport, have_voice)
+
+ # 接收音频
+ asr_component = context.components.get('asr')
+ if asr_component and hasattr(asr_component, 'asr_instance'):
+ asr_instance = asr_component.asr_instance
+ if hasattr(asr_instance, 'receive_audio'):
+ await asr_instance.receive_audio(context, audio, have_voice)
+
+ async def _resume_vad_detection(self, context: SessionContext):
+ """恢复VAD检测 - 完整迁移自resume_vad_detection"""
+ # 等待1秒后恢复VAD检测
+ await asyncio.sleep(1)
+ context.just_woken_up = False
+
+ async def start_to_chat(self, context: SessionContext, transport: TransportInterface, text: str):
+ """开始聊天 - 完整迁移自startToChat"""
+ # 检查输入是否是JSON格式(包含说话人信息)
+ speaker_name = None
+ actual_text = text
+
+ try:
+ # 尝试解析JSON格式的输入
+ if text.strip().startswith('{') and text.strip().endswith('}'):
+ data = json.loads(text)
+ if 'speaker' in data and 'content' in data:
+ speaker_name = data['speaker']
+ actual_text = data['content']
+ logger.info(f"解析到说话人信息: {speaker_name}")
+
+ # 直接使用JSON格式的文本,不解析
+ actual_text = text
+ except (json.JSONDecodeError, KeyError):
+ # 如果解析失败,继续使用原始文本
+ pass
+
+ # 保存说话人信息到上下文
+ if speaker_name:
+ context.current_speaker = speaker_name
+ else:
+ context.current_speaker = None
+
+ # 检查设备绑定
+ if context.need_bind:
+ await self._check_bind_device(context, transport)
+ return
+
+ # 如果当日的输出字数大于限定的字数
+ if context.max_output_size > 0:
+ if check_device_output_limit(
+ context.headers.get("device-id"), context.max_output_size
+ ):
+ await self._max_out_size(context, transport)
+ return
+
+ if context.is_speaking:
+ await self._handle_abort_message(context, transport)
+
+ # 首先进行意图分析,使用实际文本内容
+ from core.processors.chat_processor import ChatProcessor
+ chat_processor = ChatProcessor()
+ intent_handled = await chat_processor.handle_user_intent(context, transport, actual_text)
+
+ if intent_handled:
+ # 如果意图已被处理,不再进行聊天
+ return
+
+ # 意图未被处理,继续常规聊天流程,使用实际文本内容
+ await self._send_stt_message(context, transport, actual_text)
+
+ # 使用ChatProcessor处理聊天
+ from core.processors.chat_processor import ChatProcessor
+ chat_processor = ChatProcessor()
+ await chat_processor.handle_chat(context, transport, actual_text)
+
+ async def _no_voice_close_connect(self, context: SessionContext, transport: TransportInterface, have_voice: bool):
+ """无声音时关闭连接检测 - 完整迁移自no_voice_close_connect"""
+ if have_voice:
+ context.update_activity()
+ return
+
+ # 只有在已经初始化过时间戳的情况下才进行超时检查
+ if context.last_activity_time_ms > 0.0:
+ no_voice_time = time.time() * 1000 - context.last_activity_time_ms
+ close_connection_no_voice_time = int(
+ context.config.get("close_connection_no_voice_time", 120)
+ )
+
+ if (
+ not context.close_after_chat
+ and no_voice_time > 1000 * close_connection_no_voice_time
+ ):
+ context.close_after_chat = True
+ context.abort_requested = False
+
+ end_prompt = context.config.get("end_prompt", {})
+ if end_prompt and end_prompt.get("enable", True) is False:
+ logger.info("结束对话,无需发送结束提示语")
+ await transport.close()
+ return
+
+ prompt = end_prompt.get("prompt")
+ if not prompt:
+ prompt = "请你以```时间过得真快```未来头,用富有感情、依依不舍的话来结束这场对话吧。!"
+ await self.start_to_chat(context, transport, prompt)
+
+ async def _max_out_size(self, context: SessionContext, transport: TransportInterface):
+ """超出最大输出字数处理 - 完整迁移自max_out_size"""
+ # 播放超出最大输出字数的提示
+ context.abort_requested = False
+ text = "不好意思,我现在有点事情要忙,明天这个时候我们再聊,约好了哦!明天不见不散,拜拜!"
+ await self._send_stt_message(context, transport, text)
+
+ file_path = "config/assets/max_output_size.wav"
+ opus_packets = audio_to_data(file_path)
+
+ # 获取TTS组件并添加到队列
+ tts_component = context.components.get('tts')
+ if tts_component and hasattr(tts_component, 'tts_instance'):
+ tts_instance = tts_component.tts_instance
+ if hasattr(tts_instance, 'tts_audio_queue'):
+ from core.providers.tts.dto.dto import SentenceType
+ tts_instance.tts_audio_queue.put((SentenceType.LAST, opus_packets, text))
+
+ context.close_after_chat = True
+
+ async def _check_bind_device(self, context: SessionContext, transport: TransportInterface):
+ """检查设备绑定 - 完整迁移自check_bind_device"""
+ bind_code = context.bind_code
+
+ if bind_code:
+ # 确保bind_code是6位数字
+ if len(bind_code) != 6:
+ logger.error(f"无效的绑定码格式: {bind_code}")
+ text = "绑定码格式错误,请检查配置。"
+ await self._send_stt_message(context, transport, text)
+ return
+
+ text = f"请登录控制面板,输入{bind_code},绑定设备。"
+ await self._send_stt_message(context, transport, text)
+
+ # 获取TTS组件
+ tts_component = context.components.get('tts')
+ if not tts_component or not hasattr(tts_component, 'tts_instance'):
+ return
+
+ tts_instance = tts_component.tts_instance
+ if not hasattr(tts_instance, 'tts_audio_queue'):
+ return
+
+ # 播放提示音
+ from core.providers.tts.dto.dto import SentenceType
+ music_path = "config/assets/bind_code.wav"
+ opus_packets = audio_to_data(music_path)
+ tts_instance.tts_audio_queue.put((SentenceType.FIRST, opus_packets, text))
+
+ # 逐个播放数字
+ for i in range(6): # 确保只播放6位数字
+ try:
+ digit = bind_code[i]
+ num_path = f"config/assets/bind_code/{digit}.wav"
+ num_packets = audio_to_data(num_path)
+ tts_instance.tts_audio_queue.put((SentenceType.MIDDLE, num_packets, None))
+ except Exception as e:
+ logger.error(f"播放数字音频失败: {e}")
+ continue
+ tts_instance.tts_audio_queue.put((SentenceType.LAST, [], None))
+ else:
+ # 播放未绑定提示
+ context.abort_requested = False
+ text = f"没有找到该设备的版本信息,请正确配置 OTA地址,然后重新编译固件。"
+ await self._send_stt_message(context, transport, text)
+
+ # 获取TTS组件
+ tts_component = context.components.get('tts')
+ if tts_component and hasattr(tts_component, 'tts_instance'):
+ tts_instance = tts_component.tts_instance
+ if hasattr(tts_instance, 'tts_audio_queue'):
+ from core.providers.tts.dto.dto import SentenceType
+ music_path = "config/assets/bind_not_found.wav"
+ opus_packets = audio_to_data(music_path)
+ tts_instance.tts_audio_queue.put((SentenceType.LAST, opus_packets, text))
+
+ async def _handle_abort_message(self, context: SessionContext, transport: TransportInterface):
+ """处理中断消息"""
+ logger.info("Audio processor: Abort message received")
+ context.abort_requested = True
+
+ # 清理队列
+ await self._clear_queues(context)
+
+ # 打断客户端说话状态
+ await transport.send(json.dumps({
+ "type": "tts",
+ "state": "stop",
+ "session_id": context.session_id
+ }))
+
+ # 清理说话状态
+ context.is_speaking = False
+
+ async def _clear_queues(self, context: SessionContext):
+ """清理所有队列"""
+ # 清理TTS音频队列
+ tts_component = context.components.get('tts')
+ if tts_component and hasattr(tts_component, 'tts_instance'):
+ tts_instance = tts_component.tts_instance
+ if hasattr(tts_instance, 'tts_audio_queue'):
+ try:
+ while not tts_instance.tts_audio_queue.empty():
+ tts_instance.tts_audio_queue.get_nowait()
+ except:
+ pass
+
+ # 清理ASR音频队列
+ context.clear_audio_buffer()
+
+ async def _send_stt_message(self, context: SessionContext, transport: TransportInterface, text: str):
+ """发送STT消息"""
+ await transport.send(json.dumps({
+ "type": "stt",
+ "text": text,
+ "session_id": context.session_id
+ }))
+
+
diff --git a/main/xiaozhi-server/core/processors/audio_send_processor.py b/main/xiaozhi-server/core/processors/audio_send_processor.py
new file mode 100644
index 00000000..cc67fbcb
--- /dev/null
+++ b/main/xiaozhi-server/core/processors/audio_send_processor.py
@@ -0,0 +1,132 @@
+import json
+import time
+import asyncio
+from typing import Any, List
+from core.pipeline.message_pipeline import MessageProcessor
+from core.context.session_context import SessionContext
+from core.transport.transport_interface import TransportInterface
+from core.providers.tts.dto.dto import SentenceType
+from core.utils import textUtils
+from config.logger import setup_logging
+
+logger = setup_logging()
+
+
+class AudioSendProcessor(MessageProcessor):
+ """音频发送处理器:完整迁移sendAudioHandle.py的所有功能"""
+
+ async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
+ """这个处理器不直接处理消息,而是被其他处理器调用"""
+ return False
+
+ async def send_audio_message(self, context: SessionContext, transport: TransportInterface,
+ sentence_type: SentenceType, audios: bytes, text: str):
+ """发送音频消息 - 完整迁移自sendAudioMessage"""
+ tts_component = context.components.get('tts')
+ if not tts_component or not hasattr(tts_component, 'tts_instance'):
+ return
+
+ tts_instance = tts_component.tts_instance
+
+ if hasattr(tts_instance, 'tts_audio_first_sentence') and tts_instance.tts_audio_first_sentence:
+ logger.info(f"发送第一段语音: {text}")
+ tts_instance.tts_audio_first_sentence = False
+ await self.send_tts_message(context, transport, "start", None)
+
+ if sentence_type == SentenceType.FIRST:
+ await self.send_tts_message(context, transport, "sentence_start", text)
+
+ await self.send_audio(context, transport, audios)
+
+ # 发送句子开始消息
+ if sentence_type is not SentenceType.MIDDLE:
+ logger.info(f"发送音频消息: {sentence_type}, {text}")
+
+ # 发送结束消息(如果是最后一个文本)
+ if context.llm_finish_task and sentence_type == SentenceType.LAST:
+ await self.send_tts_message(context, transport, "stop", None)
+ context.is_speaking = False
+ if context.close_after_chat:
+ await transport.close()
+
+ async def send_audio(self, context: SessionContext, transport: TransportInterface,
+ audios: bytes, frame_duration: int = 60):
+ """发送单个opus包,支持流控 - 完整迁移自sendAudio"""
+ if audios is None or len(audios) == 0:
+ return
+
+ if isinstance(audios, bytes):
+ if context.abort_requested:
+ return
+
+ context.update_activity()
+ await transport.send(audios)
+ await asyncio.sleep(frame_duration / 1000.0)
+ elif isinstance(audios, list):
+ for audio in audios:
+ if context.abort_requested:
+ break
+ context.update_activity()
+ await transport.send(audio)
+ await asyncio.sleep(frame_duration / 1000.0)
+
+ async def send_stt_message(self, context: SessionContext, transport: TransportInterface, text: str):
+ """发送STT消息 - 完整迁移自send_stt_message"""
+ await transport.send(json.dumps({
+ "type": "stt",
+ "text": text,
+ "session_id": context.session_id
+ }))
+ logger.info(f"发送STT消息: {text}")
+
+ async def send_tts_message(self, context: SessionContext, transport: TransportInterface,
+ state: str, text: str = None):
+ """发送TTS消息 - 完整迁移自send_tts_message"""
+ message = {
+ "type": "tts",
+ "state": state,
+ "session_id": context.session_id
+ }
+ if text:
+ message["text"] = text
+
+ await transport.send(json.dumps(message))
+ logger.debug(f"发送TTS消息: state={state}, text={text}")
+
+ async def send_music_message(self, context: SessionContext, transport: TransportInterface,
+ music_path: str, text: str):
+ """发送音乐消息 - 完整迁移自send_music_message"""
+ from core.utils.util import audio_to_data
+
+ try:
+ # 获取音频数据
+ opus_packets = audio_to_data(music_path)
+ if opus_packets:
+ # 发送音乐开始消息
+ await self.send_tts_message(context, transport, "start", text)
+
+ # 发送音频数据
+ await self.send_audio(context, transport, opus_packets)
+
+ # 发送音乐结束消息
+ await self.send_tts_message(context, transport, "stop", None)
+
+ logger.info(f"发送音乐: {music_path}")
+ else:
+ logger.warning(f"无法加载音乐文件: {music_path}")
+
+ except Exception as e:
+ logger.error(f"发送音乐失败: {e}")
+
+ async def send_welcome_audio(self, context: SessionContext, transport: TransportInterface):
+ """发送欢迎音频"""
+ welcome_audio_path = context.config.get("welcome_audio_path")
+ if welcome_audio_path:
+ await self.send_music_message(context, transport, welcome_audio_path, "欢迎使用小智助手")
+
+ async def send_goodbye_audio(self, context: SessionContext, transport: TransportInterface):
+ """发送告别音频"""
+ goodbye_audio_path = context.config.get("goodbye_audio_path")
+ if goodbye_audio_path:
+ await self.send_music_message(context, transport, goodbye_audio_path, "再见,期待下次相遇")
+
diff --git a/main/xiaozhi-server/core/processors/auth_processor.py b/main/xiaozhi-server/core/processors/auth_processor.py
new file mode 100644
index 00000000..9aa5f2ca
--- /dev/null
+++ b/main/xiaozhi-server/core/processors/auth_processor.py
@@ -0,0 +1,44 @@
+from typing import Any
+from core.pipeline.message_pipeline import MessageProcessor
+from core.context.session_context import SessionContext
+from core.transport.transport_interface import TransportInterface
+from core.auth import AuthMiddleware, AuthenticationError
+from config.logger import setup_logging
+
+logger = setup_logging()
+
+
+class AuthProcessor(MessageProcessor):
+ """认证处理器:处理连接认证逻辑"""
+
+ def __init__(self):
+ self.auth_middleware = None
+
+ async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
+ """处理认证相关逻辑"""
+ # 如果已经认证,跳过
+ if context.is_authenticated:
+ return False
+
+ # 初始化认证中间件(延迟初始化)
+ if self.auth_middleware is None:
+ self.auth_middleware = AuthMiddleware(context.config)
+
+ # 检查是否为认证消息(通过headers进行认证)
+ if context.headers:
+ try:
+ await self.auth_middleware.authenticate(context.headers)
+ context.is_authenticated = True
+ logger.info(f"设备认证成功: {context.device_id}")
+ return False # 认证成功,继续处理其他消息
+ except AuthenticationError as e:
+ logger.error(f"设备认证失败: {e}")
+ # 发送认证失败消息
+ await transport.send("Authentication failed")
+ await transport.close()
+ return True # 认证失败,停止处理
+
+ # 如果没有认证信息,要求认证
+ await transport.send("Authentication required")
+ return True # 停止后续处理
+
diff --git a/main/xiaozhi-server/core/processors/chat_processor.py b/main/xiaozhi-server/core/processors/chat_processor.py
new file mode 100644
index 00000000..b3b91221
--- /dev/null
+++ b/main/xiaozhi-server/core/processors/chat_processor.py
@@ -0,0 +1,516 @@
+import json
+import uuid
+import asyncio
+from typing import Any, Dict
+from core.pipeline.message_pipeline import MessageProcessor
+from core.context.session_context import SessionContext
+from core.transport.transport_interface import TransportInterface
+from core.utils.dialogue import Message, Dialogue
+from core.utils.util import remove_punctuation_and_length
+from core.providers.tts.dto.dto import ContentType, TTSMessageDTO, SentenceType
+from plugins_func.register import Action, ActionResponse
+from config.logger import setup_logging
+
+logger = setup_logging()
+
+
+class ChatProcessor(MessageProcessor):
+ """聊天处理器:完整迁移intentHandler.py的所有功能"""
+
+ def __init__(self):
+ # 会话对话历史管理
+ self._dialogues: Dict[str, Dialogue] = {}
+
+ async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
+ """处理聊天消息"""
+ # 这个处理器不直接处理原始消息,而是被其他处理器调用
+ return False
+
+ async def handle_chat(self, context: SessionContext, transport: TransportInterface, text: str):
+ """处理聊天请求 - 完整迁移自handle_user_intent"""
+ try:
+ # 首先进行意图处理
+ intent_handled = await self.handle_user_intent(context, transport, text)
+ if intent_handled:
+ return
+
+ # 如果意图未处理,进行常规聊天
+ await self._regular_chat(context, transport, text)
+
+ except Exception as e:
+ logger.error(f"处理聊天失败: {e}")
+ await self._send_error(transport, "聊天处理失败,请重试")
+
+ async def handle_user_intent(self, context: SessionContext, transport: TransportInterface, text: str):
+ """处理用户意图 - 完整迁移自intentHandler.py"""
+ # 预处理输入文本,处理可能的JSON格式
+ try:
+ if text.strip().startswith('{') and text.strip().endswith('}'):
+ parsed_data = json.loads(text)
+ if isinstance(parsed_data, dict) and "content" in parsed_data:
+ text = parsed_data["content"] # 提取content用于意图分析
+ context.current_speaker = parsed_data.get("speaker") # 保留说话人信息
+ except (json.JSONDecodeError, TypeError):
+ pass
+
+ # 检查是否有明确的退出命令
+ _, filtered_text = remove_punctuation_and_length(text)
+ if await self._check_direct_exit(context, transport, filtered_text):
+ return True
+
+ # 检查是否是唤醒词
+ if await self._check_wakeup_words(context, transport, filtered_text):
+ return True
+
+ if context.intent_type == "function_call":
+ # 使用支持function calling的聊天方法,不再进行意图分析
+ return False
+
+ # 使用LLM进行意图分析
+ intent_result = await self._analyze_intent_with_llm(context, text)
+ if not intent_result:
+ return False
+
+ # 会话开始时生成sentence_id
+ context.sentence_id = str(uuid.uuid4().hex)
+
+ # 处理各种意图
+ return await self._process_intent_result(context, transport, intent_result, text)
+
+ def _get_dialogue(self, session_id: str) -> Dialogue:
+ """获取或创建对话历史"""
+ if session_id not in self._dialogues:
+ self._dialogues[session_id] = Dialogue()
+ return self._dialogues[session_id]
+
+ async def _get_memory_context(self, context: SessionContext, query: str) -> str:
+ """获取记忆上下文"""
+ try:
+ memory_component = context.components.get('memory')
+ if memory_component and hasattr(memory_component, 'memory_instance'):
+ memory_instance = memory_component.memory_instance
+ if hasattr(memory_instance, 'query_memory'):
+ return await memory_instance.query_memory(query)
+ except Exception as e:
+ logger.warning(f"获取记忆上下文失败: {e}")
+
+ return None
+
+ async def _generate_llm_response(self, context: SessionContext, transport: TransportInterface,
+ llm_instance, dialogue_context: list, dialogue: Dialogue):
+ """生成LLM回复"""
+ try:
+ # 初始化sentence_id并发送TTS FIRST标记(模拟原connection.py第692-700行)
+ if not context.sentence_id:
+ context.sentence_id = str(uuid.uuid4().hex)
+
+ # 发送TTS开始标记
+ await self._send_tts_first_marker(context)
+
+ # 检查是否支持流式响应
+ if hasattr(llm_instance, 'response'):
+ # 使用流式响应
+ response_generator = llm_instance.response(context.session_id, dialogue_context)
+
+ response_parts = []
+ async for response_part in self._async_generator_wrapper(response_generator):
+ if context.abort_requested:
+ break
+
+ if response_part and len(response_part) > 0:
+ response_parts.append(response_part)
+
+ # 原架构不发送流式响应给前端,直接进行TTS处理
+ # 将响应片段放入TTS队列进行语音合成
+ await self._process_response_part_for_tts(context, response_part)
+
+ # 完整回复
+ full_response = "".join(response_parts)
+ if full_response:
+ # 添加助手回复到对话历史
+ dialogue.put(Message(role="assistant", content=full_response))
+
+ # 原架构不发送response_complete给前端,只进行TTS处理
+ # 发送TTS结束标记
+ await self._finalize_tts_response(context, full_response)
+
+ logger.info(f"LLM回复完成: {full_response[:100]}...")
+
+ else:
+ logger.warning("LLM实例不支持流式响应")
+
+ except Exception as e:
+ logger.error(f"生成LLM回复失败: {e}")
+ await self._send_error(transport, "生成回复失败")
+
+ async def _async_generator_wrapper(self, generator):
+ """将同步生成器包装为异步生成器"""
+ try:
+ for item in generator:
+ yield item
+ # 让出控制权,避免阻塞事件循环
+ await asyncio.sleep(0)
+ except Exception as e:
+ logger.error(f"生成器包装失败: {e}")
+
+ async def _send_tts_first_marker(self, context: SessionContext):
+ """发送TTS开始标记"""
+ try:
+ tts_component = context.components.get('tts')
+ if not tts_component or not hasattr(tts_component, 'tts_instance'):
+ return
+
+ tts_instance = tts_component.tts_instance
+ if not tts_instance or not hasattr(tts_instance, 'tts_text_queue'):
+ return
+
+ # 发送TTS开始标记(模拟原connection.py第694-700行)
+ tts_instance.tts_text_queue.put(TTSMessageDTO(
+ sentence_id=context.sentence_id,
+ sentence_type=SentenceType.FIRST,
+ content_type=ContentType.ACTION
+ ))
+
+ except Exception as e:
+ logger.error(f"发送TTS开始标记失败: {e}")
+
+ async def _process_response_part_for_tts(self, context: SessionContext, response_part: str):
+ """处理响应片段进行TTS - 模拟原架构逻辑"""
+ try:
+ tts_component = context.components.get('tts')
+ if not tts_component or not hasattr(tts_component, 'tts_instance'):
+ return
+
+ tts_instance = tts_component.tts_instance
+ if not tts_instance or not hasattr(tts_instance, 'tts_text_queue'):
+ return
+
+ # 将响应片段放入TTS队列(模拟原connection.py第782-789行逻辑)
+ tts_instance.tts_text_queue.put(TTSMessageDTO(
+ sentence_id=context.sentence_id,
+ sentence_type=SentenceType.MIDDLE,
+ content_type=ContentType.TEXT,
+ content_detail=response_part
+ ))
+
+ except Exception as e:
+ logger.error(f"处理TTS响应片段失败: {e}")
+
+ async def _finalize_tts_response(self, context: SessionContext, full_response: str):
+ """完成TTS响应 - 发送结束标记"""
+ try:
+ tts_component = context.components.get('tts')
+ if not tts_component or not hasattr(tts_component, 'tts_instance'):
+ return
+
+ tts_instance = tts_component.tts_instance
+ if not tts_instance or not hasattr(tts_instance, 'tts_text_queue'):
+ return
+
+ # 发送TTS结束标记(模拟原speak_txt函数逻辑)
+ tts_instance.tts_text_queue.put(TTSMessageDTO(
+ sentence_id=context.sentence_id,
+ sentence_type=SentenceType.LAST,
+ content_type=ContentType.ACTION
+ ))
+
+ # 设置LLM完成标记
+ context.llm_finish_task = True
+
+ except Exception as e:
+ logger.error(f"完成TTS响应失败: {e}")
+
+ async def _trigger_tts(self, context: SessionContext, transport: TransportInterface, text: str):
+ """触发TTS语音合成 - 完整迁移自原chat方法的TTS处理"""
+ try:
+ tts_component = context.components.get('tts')
+ if not tts_component or not hasattr(tts_component, 'tts_instance'):
+ logger.warning("TTS组件未初始化")
+ return
+
+ tts_instance = tts_component.tts_instance
+
+ # 确保有sentence_id
+ if not context.sentence_id:
+ context.sentence_id = str(uuid.uuid4().hex)
+
+ logger.info(f"触发TTS合成: {text[:50]}...")
+
+ # 使用原来的TTS处理方式
+ if hasattr(tts_instance, 'tts_text_queue') and hasattr(tts_instance, 'tts_one_sentence'):
+ # 发送FIRST消息到TTS队列
+ tts_instance.tts_text_queue.put(
+ TTSMessageDTO(
+ sentence_id=context.sentence_id,
+ sentence_type=SentenceType.FIRST,
+ content_type=ContentType.ACTION,
+ )
+ )
+
+ # 合成一句话
+ tts_instance.tts_one_sentence(context, ContentType.TEXT, content_detail=text)
+
+ # 发送LAST消息到TTS队列
+ tts_instance.tts_text_queue.put(
+ TTSMessageDTO(
+ sentence_id=context.sentence_id,
+ sentence_type=SentenceType.LAST,
+ content_type=ContentType.ACTION,
+ )
+ )
+
+ logger.info("TTS合成任务已提交到队列")
+ else:
+ logger.warning("TTS实例不支持队列处理")
+
+ except Exception as e:
+ logger.error(f"TTS合成失败: {e}")
+
+ async def _send_error(self, transport: TransportInterface, error_message: str):
+ """发送错误消息"""
+ try:
+ await transport.send(json.dumps({
+ "type": "error",
+ "message": error_message
+ }))
+ except Exception as e:
+ logger.error(f"发送错误消息失败: {e}")
+
+ # === 意图处理相关方法:完整迁移自intentHandler.py ===
+
+ async def _check_direct_exit(self, context: SessionContext, transport: TransportInterface, text: str):
+ """检查是否有明确的退出命令 - 完整迁移自check_direct_exit"""
+ _, text = remove_punctuation_and_length(text)
+ cmd_exit = context.cmd_exit
+ for cmd in cmd_exit:
+ if text == cmd:
+ logger.info(f"识别到明确的退出命令: {text}")
+ await self._send_stt_message(context, transport, text)
+ await transport.close()
+ return True
+ return False
+
+ async def _check_wakeup_words(self, context: SessionContext, transport: TransportInterface, text: str):
+ """检查唤醒词 - 调用TextProcessor的方法"""
+ # 这里需要调用TextProcessor的checkWakeupWords方法
+ # 为了避免循环依赖,我们在这里实现简化版本
+ _, filtered_text = remove_punctuation_and_length(text)
+ if filtered_text in context.config.get("wakeup_words", []):
+ return True
+ return False
+
+ async def _analyze_intent_with_llm(self, context: SessionContext, text: str):
+ """使用LLM分析用户意图 - 完整迁移自analyze_intent_with_llm"""
+ intent_component = context.components.get('intent')
+ if not intent_component or not hasattr(intent_component, 'intent_instance'):
+ logger.warning("意图识别服务未初始化")
+ return None
+
+ intent_instance = intent_component.intent_instance
+
+ # 对话历史记录
+ dialogue = context.dialogue
+ if not dialogue:
+ return None
+
+ try:
+ intent_result = await intent_instance.detect_intent(context, dialogue.dialogue, text)
+ return intent_result
+ except Exception as e:
+ logger.error(f"意图识别失败: {str(e)}")
+
+ return None
+
+ async def _process_intent_result(self, context: SessionContext, transport: TransportInterface, intent_result: str, original_text: str):
+ """处理意图识别结果 - 完整迁移自process_intent_result"""
+ try:
+ # 尝试将结果解析为JSON
+ intent_data = json.loads(intent_result)
+
+ # 检查是否有function_call
+ if "function_call" in intent_data:
+ # 直接从意图识别获取了function_call
+ logger.debug(f"检测到function_call格式的意图结果: {intent_data['function_call']['name']}")
+ function_name = intent_data["function_call"]["name"]
+ if function_name == "continue_chat":
+ return False
+
+ function_args = {}
+ if "arguments" in intent_data["function_call"]:
+ function_args = intent_data["function_call"]["arguments"]
+ if function_args is None:
+ function_args = {}
+
+ # 确保参数是字符串格式的JSON
+ if isinstance(function_args, dict):
+ function_args = json.dumps(function_args)
+
+ function_call_data = {
+ "name": function_name,
+ "id": str(uuid.uuid4().hex),
+ "arguments": function_args,
+ }
+
+ await self._send_stt_message(context, transport, original_text)
+ context.abort_requested = False
+
+ # 使用executor执行函数调用和结果处理
+ await self._process_function_call(context, transport, function_call_data, original_text)
+ return True
+ return False
+ except json.JSONDecodeError as e:
+ logger.error(f"处理意图结果时出错: {e}")
+ return False
+
+ async def _process_function_call(self, context: SessionContext, transport: TransportInterface, function_call_data: dict, original_text: str):
+ """处理函数调用 - 完整迁移自process_function_call"""
+ def process_function_call():
+ # 添加用户消息到对话历史
+ dialogue = context.dialogue
+ if dialogue:
+ dialogue.put(Message(role="user", content=original_text))
+
+ # 使用统一工具处理器处理所有工具调用
+ try:
+ func_handler = context.func_handler
+ if not func_handler:
+ raise Exception("函数处理器未初始化")
+
+ loop = context.loop
+ result = asyncio.run_coroutine_threadsafe(
+ func_handler.handle_llm_function_call(context, function_call_data),
+ loop,
+ ).result()
+ except Exception as e:
+ logger.error(f"工具调用失败: {e}")
+ result = ActionResponse(
+ action=Action.ERROR, result=str(e), response=str(e)
+ )
+
+ if result:
+ function_name = function_call_data.get("name", "")
+ if result.action == Action.RESPONSE: # 直接回复前端
+ text = result.response
+ if text is not None:
+ self._speak_txt(context, text)
+ elif result.action == Action.REQLLM: # 调用函数后再请求llm生成回复
+ text = result.result
+ if dialogue:
+ dialogue.put(Message(role="tool", content=text))
+ intent_component = context.components.get('intent')
+ if intent_component and hasattr(intent_component, 'intent_instance'):
+ intent_instance = intent_component.intent_instance
+ if hasattr(intent_instance, 'replyResult'):
+ llm_result = intent_instance.replyResult(text, original_text)
+ if llm_result is None:
+ llm_result = text
+ self._speak_txt(context, llm_result)
+ elif (
+ result.action == Action.NOTFOUND
+ or result.action == Action.ERROR
+ ):
+ text = result.result
+ if text is not None:
+ self._speak_txt(context, text)
+ elif function_name != "play_music":
+ # For backward compatibility with original code
+ # 获取当前最新的文本索引
+ text = result.response
+ if text is None:
+ text = result.result
+ if text is not None:
+ self._speak_txt(context, text)
+
+ # 将函数执行放在线程池中
+ if context.executor:
+ context.executor.submit(process_function_call)
+ else:
+ # 如果没有executor,直接执行
+ process_function_call()
+
+ def _speak_txt(self, context: SessionContext, text: str):
+ """语音合成文本 - 完整迁移自speak_txt"""
+ tts_component = context.components.get('tts')
+ if not tts_component or not hasattr(tts_component, 'tts_instance'):
+ return
+
+ tts_instance = tts_component.tts_instance
+ sentence_id = context.sentence_id or str(uuid.uuid4().hex)
+
+ # 发送TTS消息队列
+ if hasattr(tts_instance, 'tts_text_queue'):
+ tts_instance.tts_text_queue.put(
+ TTSMessageDTO(
+ sentence_id=sentence_id,
+ sentence_type=SentenceType.FIRST,
+ content_type=ContentType.ACTION,
+ )
+ )
+
+ # 合成一句话
+ if hasattr(tts_instance, 'tts_one_sentence'):
+ tts_instance.tts_one_sentence(context, ContentType.TEXT, content_detail=text)
+
+ tts_instance.tts_text_queue.put(
+ TTSMessageDTO(
+ sentence_id=sentence_id,
+ sentence_type=SentenceType.LAST,
+ content_type=ContentType.ACTION,
+ )
+ )
+
+ # 添加到对话历史
+ dialogue = context.dialogue
+ if dialogue:
+ dialogue.put(Message(role="assistant", content=text))
+
+ async def _regular_chat(self, context: SessionContext, transport: TransportInterface, text: str):
+ """常规聊天处理"""
+ # 使用SessionContext的对话历史
+ dialogue = context.dialogue
+ if not dialogue:
+ from core.utils.dialogue import Dialogue
+ dialogue = Dialogue()
+ context.dialogue = dialogue
+
+ # 获取LLM组件
+ llm_component = context.components.get('llm')
+ if not llm_component:
+ await self._send_error(transport, "LLM组件未初始化")
+ return
+
+ llm_instance = getattr(llm_component, 'llm_instance', None)
+ if not llm_instance:
+ await self._send_error(transport, "LLM实例未就绪")
+ return
+
+ # 添加用户消息到对话历史
+ dialogue.put(Message(role="user", content=text))
+
+ # 原架构不发送thinking状态给前端,直接开始处理
+
+ # 获取记忆上下文
+ memory_context = await self._get_memory_context(context, text)
+
+ # 构建对话上下文
+ dialogue_context = dialogue.get_llm_dialogue_with_memory(
+ memory_context,
+ context.config.get("voiceprint", {})
+ )
+
+ # 调用LLM生成回复
+ await self._generate_llm_response(context, transport, llm_instance, dialogue_context, dialogue)
+
+ async def _send_stt_message(self, context: SessionContext, transport: TransportInterface, text: str):
+ """发送STT消息"""
+ await transport.send(json.dumps({
+ "type": "stt",
+ "text": text,
+ "session_id": context.session_id
+ }))
+
+ def cleanup_session(self, session_id: str):
+ """清理会话对话历史"""
+ if session_id in self._dialogues:
+ del self._dialogues[session_id]
+ logger.info(f"已清理会话对话历史: {session_id}")
diff --git a/main/xiaozhi-server/core/processors/hello_processor.py b/main/xiaozhi-server/core/processors/hello_processor.py
new file mode 100644
index 00000000..36b2aaa1
--- /dev/null
+++ b/main/xiaozhi-server/core/processors/hello_processor.py
@@ -0,0 +1,206 @@
+import time
+import json
+import random
+import asyncio
+from typing import Any
+from core.pipeline.message_pipeline import MessageProcessor
+from core.context.session_context import SessionContext
+from core.transport.transport_interface import TransportInterface
+from core.utils.dialogue import Message
+from core.utils.util import audio_to_data, remove_punctuation_and_length, opus_datas_to_wav_bytes
+from core.providers.tts.dto.dto import SentenceType
+from core.utils.wakeup_word import WakeupWordsConfig
+from core.providers.tools.device_mcp import (
+ MCPClient,
+ send_mcp_initialize_message,
+ send_mcp_tools_list_request,
+)
+from config.logger import setup_logging
+
+logger = setup_logging()
+
+# 唤醒词配置
+WAKEUP_CONFIG = {
+ "refresh_time": 5,
+ "words": ["你好", "你好啊", "嘿,你好", "嗨"],
+}
+
+# 创建全局的唤醒词配置管理器
+wakeup_words_config = WakeupWordsConfig()
+
+# 用于防止并发调用wakeupWordsResponse的锁
+_wakeup_response_lock = asyncio.Lock()
+
+
+class HelloProcessor(MessageProcessor):
+ """Hello消息处理器:完整迁移helloHandle.py的所有功能"""
+
+ async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
+ """处理hello类型的消息"""
+ if isinstance(message, str):
+ try:
+ msg_json = json.loads(message)
+ if isinstance(msg_json, dict) and msg_json.get("type") == "hello":
+ await self.handle_hello_message(context, transport, msg_json)
+ return True
+ except json.JSONDecodeError:
+ pass
+ return False
+
+ async def handle_hello_message(self, context: SessionContext, transport: TransportInterface, msg_json: dict):
+ """处理hello消息 - 完整迁移自handleHelloMessage"""
+ # 处理音频参数
+ audio_params = msg_json.get("audio_params")
+ if audio_params:
+ format = audio_params.get("format")
+ logger.info(f"客户端音频格式: {format}")
+ context.audio_format = format
+ if not context.welcome_msg:
+ context.welcome_msg = {}
+ context.welcome_msg["audio_params"] = audio_params
+
+ # 处理客户端特性
+ features = msg_json.get("features")
+ if features:
+ logger.info(f"客户端特性: {features}")
+ context.features = features
+ if features.get("mcp"):
+ logger.info("客户端支持MCP")
+ context.mcp_client = MCPClient()
+ # 发送初始化 - 传递transport参数
+ asyncio.create_task(send_mcp_initialize_message(context, transport))
+ # 发送mcp消息,获取tools列表 - 传递transport参数
+ asyncio.create_task(send_mcp_tools_list_request(context, transport))
+
+ # 发送欢迎消息
+ if context.welcome_msg:
+ await transport.send(json.dumps(context.welcome_msg))
+ else:
+ # 默认欢迎消息
+ welcome_msg = {
+ "type": "hello",
+ "session_id": context.session_id,
+ "version": 1,
+ "transport": "websocket"
+ }
+ await transport.send(json.dumps(welcome_msg))
+
+ async def check_wakeup_words(self, context: SessionContext, transport: TransportInterface, text: str) -> bool:
+ """检查唤醒词 - 完整迁移自checkWakeupWords"""
+ enable_wakeup_words_response_cache = context.config.get("enable_wakeup_words_response_cache", False)
+
+ # 等待tts初始化,最多等待3秒
+ tts_component = context.components.get('tts')
+ start_time = time.time()
+ while time.time() - start_time < 3:
+ if tts_component and hasattr(tts_component, 'tts_instance'):
+ break
+ await asyncio.sleep(0.1)
+ else:
+ return False
+
+ if not enable_wakeup_words_response_cache:
+ return False
+
+ _, filtered_text = remove_punctuation_and_length(text)
+ if filtered_text not in context.config.get("wakeup_words", []):
+ return False
+
+ context.just_woken_up = True
+ await self._send_stt_message(context, transport, text)
+
+ # 获取当前音色
+ tts_instance = getattr(tts_component, 'tts_instance', None) if tts_component else None
+ voice = getattr(tts_instance, "voice", "default") if tts_instance else "default"
+ if not voice:
+ voice = "default"
+
+ # 获取唤醒词回复配置
+ response = wakeup_words_config.get_wakeup_response(voice)
+ if not response or not response.get("file_path"):
+ response = {
+ "voice": "default",
+ "file_path": "config/assets/wakeup_words.wav",
+ "time": 0,
+ "text": "哈啰啊,我是小智啦,声音好听的台湾女孩一枚,超开心认识你耶,最近在忙啥,别忘了给我来点有趣的料哦,我超爱听八卦的啦",
+ }
+
+ # 获取音频数据
+ opus_packets = audio_to_data(response.get("file_path"))
+ # 播放唤醒词回复
+ context.abort_requested = False
+
+ logger.info(f"播放唤醒词回复: {response.get('text')}")
+ await self._send_audio_message(context, transport, SentenceType.FIRST, opus_packets, response.get("text"))
+ await self._send_audio_message(context, transport, SentenceType.LAST, [], None)
+
+ # 补充对话
+ if context.dialogue:
+ context.dialogue.put(Message(role="assistant", content=response.get("text")))
+
+ # 检查是否需要更新唤醒词回复
+ if time.time() - response.get("time", 0) > WAKEUP_CONFIG["refresh_time"]:
+ if not _wakeup_response_lock.locked():
+ asyncio.create_task(self._wakeup_words_response(context, transport))
+ return True
+
+ async def _wakeup_words_response(self, context: SessionContext, transport: TransportInterface):
+ """生成唤醒词回复 - 完整迁移自wakeupWordsResponse"""
+ tts_component = context.components.get('tts')
+ llm_component = context.components.get('llm')
+
+ tts_instance = getattr(tts_component, 'tts_instance', None) if tts_component else None
+ llm_instance = getattr(llm_component, 'llm_instance', None) if llm_component else None
+
+ if not tts_instance or not llm_instance or not hasattr(llm_instance, 'response_no_stream'):
+ return
+
+ try:
+ # 尝试获取锁,如果获取不到就返回
+ async with _wakeup_response_lock:
+ # 生成唤醒词回复
+ wakeup_word = random.choice(WAKEUP_CONFIG["words"])
+ question = (
+ "此刻用户正在和你说```"
+ + wakeup_word
+ + "```。\n请你根据以上用户的内容进行20-30字回复。要符合系统设置的角色情感和态度,不要像机器人一样说话。\n"
+ + "请勿对这条内容本身进行任何解释和回应,请勿返回表情符号,仅返回对用户的内容的回复。"
+ )
+
+ result = llm_instance.response_no_stream(context.config.get("prompt", ""), question)
+ if not result or len(result) == 0:
+ return
+
+ # 生成TTS音频
+ tts_result = await asyncio.to_thread(tts_instance.to_tts, result)
+ if not tts_result:
+ return
+
+ # 获取当前音色
+ voice = getattr(tts_instance, "voice", "default")
+
+ wav_bytes = opus_datas_to_wav_bytes(tts_result, sample_rate=16000)
+ file_path = wakeup_words_config.generate_file_path(voice)
+ with open(file_path, "wb") as f:
+ f.write(wav_bytes)
+ # 更新配置
+ wakeup_words_config.update_wakeup_response(voice, file_path, result)
+ except Exception as e:
+ logger.error(f"生成唤醒词回复失败: {e}")
+
+ async def _send_stt_message(self, context: SessionContext, transport: TransportInterface, text: str):
+ """发送STT消息"""
+ await transport.send(json.dumps({
+ "type": "stt",
+ "text": text,
+ "session_id": context.session_id
+ }))
+
+ async def _send_audio_message(self, context: SessionContext, transport: TransportInterface,
+ sentence_type: SentenceType, audios: bytes, text: str):
+ """发送音频消息"""
+ # 这里应该调用AudioSendProcessor
+ from core.processors.audio_send_processor import AudioSendProcessor
+ audio_send_processor = AudioSendProcessor()
+ await audio_send_processor.send_audio_message(context, transport, sentence_type, audios, text)
+
diff --git a/main/xiaozhi-server/core/processors/iot_processor.py b/main/xiaozhi-server/core/processors/iot_processor.py
new file mode 100644
index 00000000..a536ad10
--- /dev/null
+++ b/main/xiaozhi-server/core/processors/iot_processor.py
@@ -0,0 +1,123 @@
+import asyncio
+import json
+from typing import Any
+from core.pipeline.message_pipeline import MessageProcessor
+from core.context.session_context import SessionContext
+from core.transport.transport_interface import TransportInterface
+from core.providers.tools.device_iot import handleIotStatus, handleIotDescriptors
+from config.logger import setup_logging
+
+logger = setup_logging()
+
+
+class IotProcessor(MessageProcessor):
+ """IoT消息处理器:完整迁移iotMessageHandler.py的所有功能"""
+
+ async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
+ """处理iot类型的消息"""
+ if isinstance(message, str):
+ try:
+ msg_json = json.loads(message)
+ if isinstance(msg_json, dict) and msg_json.get("type") == "iot":
+ await self.handle_iot_message(context, transport, msg_json)
+ return True
+ except json.JSONDecodeError:
+ pass
+ return False
+
+ async def handle_iot_message(self, context: SessionContext, transport: TransportInterface, msg_json: dict):
+ """处理IoT消息 - 完整迁移自iotMessageHandler.py"""
+ tasks = []
+
+ # 处理设备描述符 - 完整迁移原逻辑
+ if "descriptors" in msg_json:
+ logger.debug("处理IoT设备描述符")
+ task = asyncio.create_task(
+ self._handle_iot_descriptors(context, transport, msg_json["descriptors"])
+ )
+ tasks.append(task)
+
+ # 处理设备状态 - 完整迁移原逻辑
+ if "states" in msg_json:
+ logger.debug("处理IoT设备状态")
+ task = asyncio.create_task(
+ self._handle_iot_status(context, transport, msg_json["states"])
+ )
+ tasks.append(task)
+
+ # 如果没有有效的IoT数据
+ if not tasks:
+ logger.warning("IoT消息缺少descriptors或states字段")
+ await self._send_error_response(
+ transport,
+ context.session_id,
+ "IoT消息格式错误:缺少descriptors或states字段"
+ )
+ return
+
+ # 等待所有任务完成(可选,根据原逻辑决定)
+ # await asyncio.gather(*tasks, return_exceptions=True)
+
+ async def _handle_iot_descriptors(self, context: SessionContext, transport: TransportInterface, descriptors: Any):
+ """处理IoT设备描述符 - 包装原handleIotDescriptors函数"""
+ try:
+ # 调用原有的handleIotDescriptors函数
+ # 注意:这里需要传入context而不是conn,因为handleIotDescriptors可能需要适配
+ await handleIotDescriptors(context, descriptors)
+ logger.debug("IoT设备描述符处理完成")
+
+ except Exception as e:
+ logger.error(f"处理IoT设备描述符失败: {e}", exc_info=True)
+ await self._send_error_response(
+ transport,
+ context.session_id,
+ f"IoT设备描述符处理失败: {str(e)}"
+ )
+
+ async def _handle_iot_status(self, context: SessionContext, transport: TransportInterface, states: Any):
+ """处理IoT设备状态 - 包装原handleIotStatus函数"""
+ try:
+ # 调用原有的handleIotStatus函数
+ # 注意:这里需要传入context而不是conn,因为handleIotStatus可能需要适配
+ await handleIotStatus(context, states)
+ logger.debug("IoT设备状态处理完成")
+
+ except Exception as e:
+ logger.error(f"处理IoT设备状态失败: {e}", exc_info=True)
+ await self._send_error_response(
+ transport,
+ context.session_id,
+ f"IoT设备状态处理失败: {str(e)}"
+ )
+
+ async def _send_error_response(self, transport: TransportInterface, session_id: str, message: str):
+ """发送IoT错误响应"""
+ response = {
+ "type": "iot",
+ "status": "error",
+ "message": message,
+ "session_id": session_id
+ }
+
+ try:
+ await transport.send(json.dumps(response))
+ except Exception as e:
+ logger.error(f"发送IoT错误响应失败: {e}")
+
+ async def _send_success_response(self, transport: TransportInterface, session_id: str,
+ message: str, data: dict = None):
+ """发送IoT成功响应"""
+ response = {
+ "type": "iot",
+ "status": "success",
+ "message": message,
+ "session_id": session_id
+ }
+ if data:
+ response["data"] = data
+
+ try:
+ await transport.send(json.dumps(response))
+ except Exception as e:
+ logger.error(f"发送IoT成功响应失败: {e}")
+
diff --git a/main/xiaozhi-server/core/processors/listen_processor.py b/main/xiaozhi-server/core/processors/listen_processor.py
new file mode 100644
index 00000000..fd947b7d
--- /dev/null
+++ b/main/xiaozhi-server/core/processors/listen_processor.py
@@ -0,0 +1,129 @@
+import time
+import json
+from typing import Any
+from core.pipeline.message_pipeline import MessageProcessor
+from core.context.session_context import SessionContext
+from core.transport.transport_interface import TransportInterface
+from core.utils.util import remove_punctuation_and_length
+from config.logger import setup_logging
+
+logger = setup_logging()
+
+
+class ListenProcessor(MessageProcessor):
+ """Listen消息处理器:完整迁移listenMessageHandler.py的所有功能"""
+
+ async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
+ """处理listen类型的消息"""
+ if isinstance(message, str):
+ try:
+ msg_json = json.loads(message)
+ if isinstance(msg_json, dict) and msg_json.get("type") == "listen":
+ await self.handle_listen_message(context, transport, msg_json)
+ return True
+ except json.JSONDecodeError:
+ pass
+ return False
+
+ async def handle_listen_message(self, context: SessionContext, transport: TransportInterface, msg_json: dict):
+ """处理listen消息 - 完整迁移自listenMessageHandler.py"""
+ # 设置拾音模式
+ if "mode" in msg_json:
+ context.listen_mode = msg_json["mode"]
+ logger.debug(f"客户端拾音模式:{context.listen_mode}")
+
+ # 处理不同的状态
+ state = msg_json.get("state")
+
+ if state == "start":
+ # 开始监听语音
+ context.client_have_voice = True
+ context.client_voice_stop = False
+ logger.debug("开始语音监听")
+
+ elif state == "stop":
+ # 停止监听语音
+ context.client_have_voice = True
+ context.client_voice_stop = True
+ # 如果有音频数据,处理最后的音频
+ if len(context.asr_audio) > 0:
+ await self._handle_audio_message(context, transport, b"")
+ logger.debug("停止语音监听")
+
+ elif state == "detect":
+ # 检测到文本输入
+ context.client_have_voice = False
+ context.asr_audio.clear()
+
+ if "text" in msg_json:
+ context.update_activity()
+ original_text = msg_json["text"] # 保留原始文本
+ filtered_len, filtered_text = remove_punctuation_and_length(original_text)
+
+ # 识别是否是唤醒词
+ is_wakeup_words = filtered_text in context.config.get("wakeup_words", [])
+ # 是否开启唤醒词回复
+ enable_greeting = context.config.get("enable_greeting", True)
+
+ if is_wakeup_words and not enable_greeting:
+ # 如果是唤醒词,且关闭了唤醒词回复,就不用回答
+ await self._send_stt_message(context, transport, original_text)
+ await self._send_tts_message(context, transport, "stop", None)
+ context.is_speaking = False
+
+ elif is_wakeup_words:
+ # 处理唤醒词
+ context.just_woken_up = True
+ # 上报纯文字数据(复用ASR上报功能,但不提供音频数据)
+ await self._enqueue_asr_report(context, "嘿,你好呀", [])
+ await self._start_to_chat(context, transport, "嘿,你好呀")
+
+ else:
+ # 处理普通文本
+ # 上报纯文字数据(复用ASR上报功能,但不提供音频数据)
+ await self._enqueue_asr_report(context, original_text, [])
+ # 否则需要LLM对文字内容进行答复
+ await self._start_to_chat(context, transport, original_text)
+
+ async def _handle_audio_message(self, context: SessionContext, transport: TransportInterface, audio: bytes):
+ """处理音频消息 - 调用AudioReceiveProcessor"""
+ # 这里应该调用AudioReceiveProcessor来处理音频
+ from core.processors.audio_receive_processor import AudioReceiveProcessor
+ audio_processor = AudioReceiveProcessor()
+ await audio_processor.handle_audio_message(context, transport, audio)
+
+ async def _send_stt_message(self, context: SessionContext, transport: TransportInterface, text: str):
+ """发送STT消息"""
+ await transport.send(json.dumps({
+ "type": "stt",
+ "text": text,
+ "session_id": context.session_id
+ }))
+ logger.info(f"发送STT消息: {text}")
+
+ async def _send_tts_message(self, context: SessionContext, transport: TransportInterface, state: str, text: str = None):
+ """发送TTS消息"""
+ message = {
+ "type": "tts",
+ "state": state,
+ "session_id": context.session_id
+ }
+ if text:
+ message["text"] = text
+
+ await transport.send(json.dumps(message))
+ logger.debug(f"发送TTS消息: state={state}, text={text}")
+
+ async def _enqueue_asr_report(self, context: SessionContext, text: str, audio_data: list):
+ """ASR上报队列"""
+ if context.report_asr_enable:
+ from core.processors.report_processor import ReportProcessor
+ report_processor = ReportProcessor()
+ report_processor.enqueue_asr_report(context, text, audio_data)
+
+ async def _start_to_chat(self, context: SessionContext, transport: TransportInterface, text: str):
+ """开始聊天 - 调用ChatProcessor"""
+ from core.processors.chat_processor import ChatProcessor
+ chat_processor = ChatProcessor()
+ await chat_processor.handle_chat(context, transport, text)
+
diff --git a/main/xiaozhi-server/core/processors/mcp_processor.py b/main/xiaozhi-server/core/processors/mcp_processor.py
new file mode 100644
index 00000000..92379b5b
--- /dev/null
+++ b/main/xiaozhi-server/core/processors/mcp_processor.py
@@ -0,0 +1,90 @@
+import asyncio
+import json
+from typing import Any
+from core.pipeline.message_pipeline import MessageProcessor
+from core.context.session_context import SessionContext
+from core.transport.transport_interface import TransportInterface
+from core.providers.tools.device_mcp import handle_mcp_message
+from config.logger import setup_logging
+
+logger = setup_logging()
+
+
+class McpProcessor(MessageProcessor):
+ """MCP消息处理器:完整迁移mcpMessageHandler.py的所有功能"""
+
+ async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
+ """处理mcp类型的消息"""
+ if isinstance(message, str):
+ try:
+ msg_json = json.loads(message)
+ if isinstance(msg_json, dict) and msg_json.get("type") == "mcp":
+ await self.handle_mcp_message(context, transport, msg_json)
+ return True
+ except json.JSONDecodeError:
+ pass
+ return False
+
+ async def handle_mcp_message(self, context: SessionContext, transport: TransportInterface, msg_json: dict):
+ """处理MCP消息 - 完整迁移自mcpMessageHandler.py"""
+ if "payload" in msg_json:
+ # 检查MCP客户端是否存在
+ if not context.mcp_client:
+ logger.warning("MCP客户端未初始化,无法处理MCP消息")
+ await self._send_error_response(transport, context.session_id, "MCP客户端未初始化")
+ return
+
+ # 创建异步任务处理MCP消息 - 完整迁移原逻辑
+ asyncio.create_task(
+ self._handle_mcp_payload(context, transport, msg_json["payload"])
+ )
+ else:
+ logger.warning("MCP消息缺少payload字段")
+ await self._send_error_response(transport, context.session_id, "MCP消息格式错误:缺少payload")
+
+ async def _handle_mcp_payload(self, context: SessionContext, transport: TransportInterface, payload: dict):
+ """处理MCP payload - 包装原handle_mcp_message函数"""
+ try:
+ # 调用原有的handle_mcp_message函数
+ # 注意:这里需要传入context而不是conn,因为handle_mcp_message可能需要适配
+ await handle_mcp_message(context, context.mcp_client, payload, transport)
+ logger.debug("MCP消息处理完成")
+
+ except Exception as e:
+ logger.error(f"处理MCP消息失败: {e}", exc_info=True)
+ await self._send_error_response(
+ transport,
+ context.session_id,
+ f"MCP消息处理失败: {str(e)}"
+ )
+
+ async def _send_error_response(self, transport: TransportInterface, session_id: str, message: str):
+ """发送MCP错误响应"""
+ response = {
+ "type": "mcp",
+ "status": "error",
+ "message": message,
+ "session_id": session_id
+ }
+
+ try:
+ await transport.send(json.dumps(response))
+ except Exception as e:
+ logger.error(f"发送MCP错误响应失败: {e}")
+
+ async def _send_success_response(self, transport: TransportInterface, session_id: str,
+ message: str, data: dict = None):
+ """发送MCP成功响应"""
+ response = {
+ "type": "mcp",
+ "status": "success",
+ "message": message,
+ "session_id": session_id
+ }
+ if data:
+ response["data"] = data
+
+ try:
+ await transport.send(json.dumps(response))
+ except Exception as e:
+ logger.error(f"发送MCP成功响应失败: {e}")
diff --git a/main/xiaozhi-server/core/processors/message_router.py b/main/xiaozhi-server/core/processors/message_router.py
new file mode 100644
index 00000000..18cff211
--- /dev/null
+++ b/main/xiaozhi-server/core/processors/message_router.py
@@ -0,0 +1,114 @@
+import json
+from typing import Any, List
+from core.pipeline.message_pipeline import MessageProcessor
+from core.context.session_context import SessionContext
+from core.transport.transport_interface import TransportInterface
+from core.processors.hello_processor import HelloProcessor
+from core.processors.listen_processor import ListenProcessor
+from core.processors.audio_receive_processor import AudioReceiveProcessor
+from core.processors.auth_processor import AuthProcessor
+from core.processors.timeout_processor import TimeoutProcessor
+from core.processors.server_processor import ServerProcessor
+from core.processors.mcp_processor import McpProcessor
+from core.processors.iot_processor import IotProcessor
+from core.processors.abort_processor import AbortProcessor
+from core.processors.text_processor import TextProcessor
+from config.logger import setup_logging
+
+logger = setup_logging()
+
+
+class MessageRouter(MessageProcessor):
+ """
+ 消息路由器:协调所有独立的processor
+ 按功能职责分离,避免耦合,每个processor专注单一职责
+ """
+
+ def __init__(self):
+ # 初始化所有独立的processor
+ self.auth_processor = AuthProcessor()
+ self.timeout_processor = TimeoutProcessor()
+ self.abort_processor = AbortProcessor()
+ self.hello_processor = HelloProcessor()
+ self.listen_processor = ListenProcessor()
+ self.server_processor = ServerProcessor()
+ self.mcp_processor = McpProcessor()
+ self.iot_processor = IotProcessor()
+ self.audio_receive_processor = AudioReceiveProcessor()
+ self.text_processor = TextProcessor()
+
+ # 按优先级排序的processor列表
+ self.processors: List[MessageProcessor] = [
+ self.timeout_processor, # 首先检查超时
+ self.auth_processor, # 然后检查认证
+ self.abort_processor, # 中断消息
+ self.hello_processor, # hello消息
+ self.listen_processor, # listen消息
+ self.server_processor, # 服务器消息
+ self.mcp_processor, # MCP消息
+ self.iot_processor, # IoT消息
+ self.audio_receive_processor, # 音频消息
+ self.text_processor, # 纯文本消息(放在最后,作为兜底处理)
+ ]
+
+ async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
+ """
+ 路由消息到合适的processor
+ 每个processor专注处理自己的消息类型,避免耦合
+ """
+ # 更新活动时间
+ context.update_activity()
+
+ # 按优先级顺序尝试每个processor
+ for processor in self.processors:
+ try:
+ if await processor.process(context, transport, message):
+ # 消息已被处理,记录日志并返回
+ logger.debug(f"消息被 {processor.__class__.__name__} 处理")
+ return True
+ except Exception as e:
+ logger.error(f"{processor.__class__.__name__} 处理消息时出错: {e}", exc_info=True)
+ continue
+
+ # 如果没有processor处理该消息,记录警告
+ if isinstance(message, str):
+ try:
+ msg_json = json.loads(message)
+ msg_type = msg_json.get("type", "unknown") if isinstance(msg_json, dict) else "non-dict"
+ logger.warning(f"未处理的消息类型: {msg_type}, 内容: {message[:100]}...")
+ except json.JSONDecodeError:
+ logger.warning(f"未处理的非JSON消息: {message[:100]}...")
+ elif isinstance(message, bytes):
+ logger.warning(f"未处理的二进制消息,大小: {len(message)} bytes")
+ else:
+ logger.warning(f"未处理的消息类型: {type(message)}")
+
+ return False
+
+ def add_processor(self, processor: MessageProcessor, priority: int = None):
+ """
+ 添加新的processor
+ priority: 优先级,数字越小优先级越高,None表示添加到末尾
+ """
+ if priority is None:
+ self.processors.append(processor)
+ else:
+ self.processors.insert(priority, processor)
+ logger.info(f"添加processor: {processor.__class__.__name__}")
+
+ def remove_processor(self, processor_class):
+ """移除指定类型的processor"""
+ self.processors = [p for p in self.processors if not isinstance(p, processor_class)]
+ logger.info(f"移除processor: {processor_class.__name__}")
+
+ def get_processor(self, processor_class):
+ """获取指定类型的processor"""
+ for processor in self.processors:
+ if isinstance(processor, processor_class):
+ return processor
+ return None
+
+ def list_processors(self) -> List[str]:
+ """列出所有processor的名称"""
+ return [processor.__class__.__name__ for processor in self.processors]
+
diff --git a/main/xiaozhi-server/core/processors/report_processor.py b/main/xiaozhi-server/core/processors/report_processor.py
new file mode 100644
index 00000000..daa6e391
--- /dev/null
+++ b/main/xiaozhi-server/core/processors/report_processor.py
@@ -0,0 +1,176 @@
+import time
+import queue
+import threading
+from typing import Any, List
+from core.pipeline.message_pipeline import MessageProcessor
+from core.context.session_context import SessionContext
+from core.transport.transport_interface import TransportInterface
+from config.manage_api_client import report as manage_report
+from config.logger import setup_logging
+
+logger = setup_logging()
+
+
+class ReportProcessor(MessageProcessor):
+ """上报处理器:完整迁移reportHandle.py的所有功能"""
+
+ async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
+ """这个处理器不直接处理消息,而是被其他处理器调用"""
+ return False
+
+ def enqueue_asr_report(self, context: SessionContext, text: str, audio_data: List[bytes]):
+ """ASR上报队列 - 完整迁移自enqueue_asr_report"""
+ if not context.report_asr_enable:
+ return
+
+ report_time = int(time.time())
+
+ # 将上报任务放入队列
+ context.report_queue.put({
+ "type": 1, # 用户类型
+ "text": text,
+ "audio_data": audio_data,
+ "report_time": report_time
+ })
+
+ # 确保上报线程已启动
+ self._ensure_report_thread(context)
+
+ def enqueue_tts_report(self, context: SessionContext, text: str, opus_data: bytes):
+ """TTS上报队列 - 完整迁移自enqueue_tts_report"""
+ if not context.report_tts_enable:
+ return
+
+ report_time = int(time.time())
+
+ # 将上报任务放入队列
+ context.report_queue.put({
+ "type": 2, # 智能体类型
+ "text": text,
+ "audio_data": opus_data,
+ "report_time": report_time
+ })
+
+ # 确保上报线程已启动
+ self._ensure_report_thread(context)
+
+ def _ensure_report_thread(self, context: SessionContext):
+ """确保上报线程已启动"""
+ if context.report_thread is None or not context.report_thread.is_alive():
+ context.report_thread = threading.Thread(
+ target=self._report_worker,
+ args=(context,),
+ daemon=True
+ )
+ context.report_thread.start()
+ logger.info(f"上报线程已启动: {context.session_id}")
+
+ def _report_worker(self, context: SessionContext):
+ """上报工作线程 - 完整迁移自ConnectionHandler中的上报逻辑"""
+ logger.info(f"上报工作线程启动: {context.session_id}")
+
+ while not context.stop_event.is_set():
+ try:
+ # 从队列获取上报任务
+ report_task = context.report_queue.get(timeout=1)
+
+ # 执行上报
+ self._execute_report(context, report_task)
+
+ except queue.Empty:
+ continue
+ except Exception as e:
+ logger.error(f"上报工作线程异常: {e}")
+
+ logger.info(f"上报工作线程退出: {context.session_id}")
+
+ def _execute_report(self, context: SessionContext, report_task: dict):
+ """执行聊天记录上报操作 - 完整迁移自report函数"""
+ try:
+ report_type = report_task["type"]
+ text = report_task["text"]
+ audio_data = report_task["audio_data"]
+ report_time = report_task["report_time"]
+
+ # 处理音频数据
+ processed_audio = None
+ if audio_data:
+ if isinstance(audio_data, list):
+ # ASR音频数据(多个音频片段)
+ processed_audio = self._process_asr_audio(audio_data)
+ elif isinstance(audio_data, bytes):
+ # TTS音频数据(opus格式)
+ processed_audio = self._opus_to_wav(audio_data)
+
+ # 执行上报
+ manage_report(
+ mac_address=context.device_id,
+ session_id=context.session_id,
+ chat_type=report_type,
+ content=text,
+ audio=processed_audio,
+ report_time=report_time,
+ )
+
+ logger.debug(f"上报成功: type={report_type}, text={text[:50]}...")
+
+ except Exception as e:
+ logger.error(f"聊天记录上报失败: {e}")
+
+ def _process_asr_audio(self, audio_data_list: List[bytes]) -> bytes:
+ """处理ASR音频数据"""
+ try:
+ # 将多个音频片段合并
+ combined_audio = b''.join(audio_data_list)
+ return combined_audio
+ except Exception as e:
+ logger.error(f"处理ASR音频数据失败: {e}")
+ return b''
+
+ def _opus_to_wav(self, opus_data: bytes) -> bytes:
+ """将Opus数据转换为WAV格式的字节流 - 完整迁移自opus_to_wav"""
+ try:
+ import opuslib_next
+ import io
+ import wave
+
+ # Opus解码器配置
+ sample_rate = 16000
+ channels = 1
+
+ # 创建Opus解码器
+ decoder = opuslib_next.Decoder(sample_rate, channels)
+
+ # 解码Opus数据
+ pcm_data = decoder.decode(opus_data, frame_size=960)
+
+ # 创建WAV文件
+ wav_buffer = io.BytesIO()
+ with wave.open(wav_buffer, 'wb') as wav_file:
+ wav_file.setnchannels(channels)
+ wav_file.setsampwidth(2) # 16-bit
+ wav_file.setframerate(sample_rate)
+ wav_file.writeframes(pcm_data)
+
+ return wav_buffer.getvalue()
+
+ except Exception as e:
+ logger.error(f"Opus转WAV失败: {e}")
+ return b''
+
+ def cleanup_session(self, context: SessionContext):
+ """清理会话上报资源"""
+ # 停止上报线程
+ if context.report_thread and context.report_thread.is_alive():
+ context.stop_event.set()
+ context.report_thread.join(timeout=5)
+
+ # 清理上报队列
+ try:
+ while not context.report_queue.empty():
+ context.report_queue.get_nowait()
+ except queue.Empty:
+ pass
+
+ logger.info(f"上报资源清理完成: {context.session_id}")
+
diff --git a/main/xiaozhi-server/core/processors/server_processor.py b/main/xiaozhi-server/core/processors/server_processor.py
new file mode 100644
index 00000000..cf55da8e
--- /dev/null
+++ b/main/xiaozhi-server/core/processors/server_processor.py
@@ -0,0 +1,151 @@
+import json
+from typing import Any
+from core.pipeline.message_pipeline import MessageProcessor
+from core.context.session_context import SessionContext
+from core.transport.transport_interface import TransportInterface
+from config.logger import setup_logging
+
+logger = setup_logging()
+
+
+class ServerProcessor(MessageProcessor):
+ """服务器消息处理器:完整迁移serverMessageHandler.py的所有功能"""
+
+ async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
+ """处理server类型的消息"""
+ if isinstance(message, str):
+ try:
+ msg_json = json.loads(message)
+ if isinstance(msg_json, dict) and msg_json.get("type") == "server":
+ await self.handle_server_message(context, transport, msg_json)
+ return True
+ except json.JSONDecodeError:
+ pass
+ return False
+
+ async def handle_server_message(self, context: SessionContext, transport: TransportInterface, msg_json: dict):
+ """处理server消息 - 完整迁移自serverMessageHandler.py"""
+ # 如果配置是从API读取的,则需要验证secret
+ if not context.read_config_from_api:
+ return
+
+ # 获取post请求的secret
+ post_secret = msg_json.get("content", {}).get("secret", "")
+ secret = context.config.get("manager-api", {}).get("secret", "")
+
+ # 如果secret不匹配,则返回
+ if post_secret != secret:
+ await self._send_error_response(
+ transport,
+ context.session_id,
+ "服务器密钥验证失败"
+ )
+ return
+
+ # 处理不同的action
+ action = msg_json.get("action")
+
+ if action == "update_config":
+ await self._handle_update_config(context, transport, msg_json)
+ elif action == "restart":
+ await self._handle_restart(context, transport, msg_json)
+ else:
+ await self._send_error_response(
+ transport,
+ context.session_id,
+ f"未知的服务器操作: {action}"
+ )
+
+ async def _handle_update_config(self, context: SessionContext, transport: TransportInterface, msg_json: dict):
+ """处理配置更新 - 完整迁移自update_config逻辑"""
+ try:
+ # 检查是否有服务器实例
+ if not context.server:
+ await self._send_error_response(
+ transport,
+ context.session_id,
+ "无法获取服务器实例",
+ {"action": "update_config"}
+ )
+ return
+
+ # 更新WebSocketServer的配置
+ if not await context.server.update_config():
+ await self._send_error_response(
+ transport,
+ context.session_id,
+ "更新服务器配置失败",
+ {"action": "update_config"}
+ )
+ return
+
+ # 发送成功响应
+ await self._send_success_response(
+ transport,
+ context.session_id,
+ "配置更新成功",
+ {"action": "update_config"}
+ )
+
+ except Exception as e:
+ logger.error(f"更新配置失败: {str(e)}")
+ await self._send_error_response(
+ transport,
+ context.session_id,
+ f"更新配置失败: {str(e)}",
+ {"action": "update_config"}
+ )
+
+ async def _handle_restart(self, context: SessionContext, transport: TransportInterface, msg_json: dict):
+ """处理服务器重启 - 完整迁移自handle_restart逻辑"""
+ try:
+ # 这里应该调用context的handle_restart方法
+ if hasattr(context, 'handle_restart'):
+ await context.handle_restart(msg_json)
+ else:
+ logger.warning("SessionContext没有handle_restart方法")
+ await self._send_error_response(
+ transport,
+ context.session_id,
+ "重启功能暂不可用",
+ {"action": "restart"}
+ )
+ except Exception as e:
+ logger.error(f"处理重启请求失败: {str(e)}")
+ await self._send_error_response(
+ transport,
+ context.session_id,
+ f"重启失败: {str(e)}",
+ {"action": "restart"}
+ )
+
+ async def _send_success_response(self, transport: TransportInterface, session_id: str,
+ message: str, content: dict = None):
+ """发送成功响应"""
+ response = {
+ "type": "server",
+ "status": "success",
+ "message": message,
+ "session_id": session_id
+ }
+ if content:
+ response["content"] = content
+
+ await transport.send(json.dumps(response))
+ logger.info(f"服务器操作成功: {message}")
+
+ async def _send_error_response(self, transport: TransportInterface, session_id: str,
+ message: str, content: dict = None):
+ """发送错误响应"""
+ response = {
+ "type": "server",
+ "status": "error",
+ "message": message,
+ "session_id": session_id
+ }
+ if content:
+ response["content"] = content
+
+ await transport.send(json.dumps(response))
+ logger.error(f"服务器操作失败: {message}")
+
diff --git a/main/xiaozhi-server/core/processors/text_processor.py b/main/xiaozhi-server/core/processors/text_processor.py
new file mode 100644
index 00000000..d0d5e390
--- /dev/null
+++ b/main/xiaozhi-server/core/processors/text_processor.py
@@ -0,0 +1,54 @@
+import json
+from typing import Any
+from core.pipeline.message_pipeline import MessageProcessor
+from core.context.session_context import SessionContext
+from core.transport.transport_interface import TransportInterface
+from config.logger import setup_logging
+
+logger = setup_logging()
+
+
+class TextProcessor(MessageProcessor):
+ """
+ 纯文本消息处理器:处理非JSON格式的文本消息
+ 这是新架构中缺失的重要组件,用于处理直接发送的文本聊天内容
+ """
+
+ async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
+ """处理纯文本消息"""
+ if isinstance(message, str):
+ try:
+ # 尝试解析为JSON,如果成功则不是纯文本消息
+ json.loads(message)
+ return False # JSON消息由其他processor处理
+ except json.JSONDecodeError:
+ # 确实是纯文本消息,进行聊天处理
+ await self.handle_text_message(context, transport, message)
+ return True
+ return False
+
+ async def handle_text_message(self, context: SessionContext, transport: TransportInterface, text: str):
+ """处理纯文本消息 - 直接调用ChatProcessor进行聊天"""
+ try:
+ # 记录收到纯文本消息
+ logger.info(f"收到纯文本消息: {text[:100]}...")
+
+ # 使用ChatProcessor处理聊天
+ from core.processors.chat_processor import ChatProcessor
+ chat_processor = ChatProcessor()
+ await chat_processor.handle_chat(context, transport, text)
+
+ except Exception as e:
+ logger.error(f"处理纯文本消息失败: {e}")
+ # 发送错误响应
+ await self._send_error_response(transport, "文本处理失败,请重试")
+
+ async def _send_error_response(self, transport: TransportInterface, error_message: str):
+ """发送错误响应"""
+ try:
+ await transport.send(json.dumps({
+ "type": "error",
+ "message": error_message
+ }))
+ except Exception as e:
+ logger.error(f"发送错误响应失败: {e}")
diff --git a/main/xiaozhi-server/core/processors/timeout_processor.py b/main/xiaozhi-server/core/processors/timeout_processor.py
new file mode 100644
index 00000000..006f481b
--- /dev/null
+++ b/main/xiaozhi-server/core/processors/timeout_processor.py
@@ -0,0 +1,41 @@
+import json
+from typing import Any
+from core.pipeline.message_pipeline import MessageProcessor
+from core.context.session_context import SessionContext
+from core.transport.transport_interface import TransportInterface
+from config.logger import setup_logging
+
+logger = setup_logging()
+
+
+class TimeoutProcessor(MessageProcessor):
+ """超时检查处理器:检查会话是否超时"""
+
+ async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
+ """检查会话超时"""
+ # 更新活动时间(在其他处理器中已更新,这里只检查)
+
+ # 获取超时配置
+ timeout_seconds = context.config.get("close_connection_no_voice_time", 120)
+
+ # 检查是否超时
+ if context.is_timeout(timeout_seconds):
+ logger.info(f"会话超时,准备关闭连接: {context.session_id}")
+
+ # 发送超时通知
+ timeout_msg = {
+ "type": "timeout",
+ "message": "连接超时,即将关闭",
+ "session_id": context.session_id
+ }
+
+ try:
+ await transport.send(json.dumps(timeout_msg))
+ await transport.close()
+ except Exception as e:
+ logger.error(f"发送超时消息失败: {e}")
+
+ return True # 消息已处理,停止后续处理
+
+ return False # 未超时,继续处理
+
diff --git a/main/xiaozhi-server/core/protocols/mqtt_connection.py b/main/xiaozhi-server/core/protocols/mqtt_connection.py
new file mode 100644
index 00000000..efb06eda
--- /dev/null
+++ b/main/xiaozhi-server/core/protocols/mqtt_connection.py
@@ -0,0 +1,279 @@
+import asyncio
+import json
+import time
+import uuid
+from typing import Dict, Any, Optional, Callable
+from config.logger import setup_logging
+
+logger = setup_logging()
+
+
+class MQTTConnection:
+ """
+ MQTT连接处理类:管理单个MQTT客户端连接
+ 处理MQTT协议消息和会话管理
+ """
+
+ def __init__(self, socket, connection_id: int, mqtt_server):
+ self.socket = socket
+ self.connection_id = connection_id
+ self.mqtt_server = mqtt_server
+
+ # 连接信息
+ self.client_id = None
+ self.device_id = None
+ self.username = None
+ self.password = None
+ self.session_id = None
+
+ # 协议状态
+ self.is_connected_flag = False
+ self.keep_alive_interval = 0
+ self.last_activity = time.time()
+
+ # 消息处理
+ self.message_callback = None
+ self.reply_topic = None
+
+ # UDP相关
+ self.udp_config = None
+
+ # 任务管理
+ self.keep_alive_task = None
+ self._closed = False
+
+ # 创建MQTT协议处理器
+ from core.protocols.mqtt_protocol import MQTTProtocol
+ self.protocol = MQTTProtocol(socket)
+ self._setup_protocol_handlers()
+
+ def _setup_protocol_handlers(self):
+ """设置协议事件处理"""
+ self.protocol.on('connect', self._handle_connect)
+ self.protocol.on('publish', self._handle_publish)
+ self.protocol.on('subscribe', self._handle_subscribe)
+ self.protocol.on('disconnect', self._handle_disconnect)
+ self.protocol.on('close', self._handle_close)
+ self.protocol.on('error', self._handle_error)
+
+ async def _handle_connect(self, connect_data: Dict[str, Any]):
+ """处理CONNECT消息"""
+ try:
+ self.client_id = connect_data['clientId']
+ self.username = connect_data.get('username')
+ self.password = connect_data.get('password')
+ self.keep_alive_interval = connect_data.get('keepAlive', 0) * 1000 # 转换为毫秒
+
+ logger.info(f"MQTT客户端连接: {self.client_id}")
+
+ # 解析客户端ID获取设备信息
+ if not self._parse_client_id():
+ await self.protocol.send_connack(1) # 连接被拒绝
+ await self.close()
+ return
+
+ # 生成会话ID
+ self.session_id = str(uuid.uuid4())
+
+ # 设置回复主题
+ self.reply_topic = f"devices/p2p/{self.device_id.replace(':', '_')}"
+
+ # 发送连接确认
+ await self.protocol.send_connack(0) # 连接接受
+ self.is_connected_flag = True
+
+ # 启动心跳检查
+ if self.keep_alive_interval > 0:
+ self.keep_alive_task = asyncio.create_task(self._keep_alive_check())
+
+ # 通知服务器新连接
+ await self.mqtt_server.on_client_connected(self)
+
+ except Exception as e:
+ logger.error(f"处理CONNECT消息失败: {e}")
+ await self.close()
+
+ def _parse_client_id(self) -> bool:
+ """解析客户端ID获取设备信息"""
+ try:
+ # 支持格式: GID_test@@@mac_address@@@uuid 或 GID_test@@@mac_address
+ parts = self.client_id.split('@@@')
+
+ if len(parts) >= 2:
+ self.group_id = parts[0]
+ # 将下划线替换为冒号格式的MAC地址
+ self.device_id = parts[1].replace('_', ':')
+
+ if len(parts) >= 3:
+ self.uuid = parts[2]
+
+ return True
+ else:
+ logger.error(f"无效的客户端ID格式: {self.client_id}")
+ return False
+
+ except Exception as e:
+ logger.error(f"解析客户端ID失败: {e}")
+ return False
+
+ async def _handle_publish(self, publish_data: Dict[str, Any]):
+ """处理PUBLISH消息"""
+ try:
+ topic = publish_data['topic']
+ payload = publish_data['payload']
+
+ logger.debug(f"收到MQTT发布消息: topic={topic}, payload={payload}")
+
+ # 更新活动时间
+ self.last_activity = time.time()
+
+ # 解析JSON消息
+ try:
+ message_data = json.loads(payload)
+
+ # 处理不同类型的消息
+ if message_data.get('type') == 'hello':
+ await self._handle_hello_message(message_data)
+ else:
+ # 其他消息通过回调处理
+ if self.message_callback:
+ self.message_callback(topic, payload)
+
+ except json.JSONDecodeError:
+ logger.error(f"MQTT消息JSON解析失败: {payload}")
+
+ except Exception as e:
+ logger.error(f"处理PUBLISH消息失败: {e}")
+
+ async def _handle_hello_message(self, message_data: Dict[str, Any]):
+ """处理hello消息,初始化UDP配置"""
+ try:
+ # 生成UDP加密配置
+ import os
+
+ self.udp_config = {
+ 'key': os.urandom(16),
+ 'encryption': 'aes-128-ctr',
+ 'server': self.mqtt_server.public_ip,
+ 'port': self.mqtt_server.udp_port
+ }
+
+ # 构造hello回复
+ hello_reply = {
+ 'type': 'hello',
+ 'version': message_data.get('version', 3),
+ 'session_id': self.session_id,
+ 'transport': 'udp',
+ 'udp': {
+ 'server': self.udp_config['server'],
+ 'port': self.udp_config['port'],
+ 'encryption': self.udp_config['encryption'],
+ 'key': self.udp_config['key'].hex(),
+ 'nonce': '00' * 16 # 临时nonce
+ },
+ 'audio_params': message_data.get('audio_params', {})
+ }
+
+ # 发送回复
+ await self.send_message(self.reply_topic, json.dumps(hello_reply))
+
+ logger.info(f"MQTT Hello消息处理完成: {self.client_id}")
+
+ except Exception as e:
+ logger.error(f"处理hello消息失败: {e}")
+
+ async def _handle_subscribe(self, subscribe_data: Dict[str, Any]):
+ """处理SUBSCRIBE消息"""
+ try:
+ topic = subscribe_data['topic']
+ packet_id = subscribe_data['packetId']
+
+ logger.debug(f"客户端订阅主题: {topic}")
+
+ # 发送订阅确认
+ await self.protocol.send_suback(packet_id, 0)
+
+ except Exception as e:
+ logger.error(f"处理SUBSCRIBE消息失败: {e}")
+
+ async def _handle_disconnect(self):
+ """处理DISCONNECT消息"""
+ logger.info(f"客户端主动断开连接: {self.client_id}")
+ await self.close()
+
+ async def _handle_close(self):
+ """处理连接关闭"""
+ logger.info(f"MQTT连接关闭: {self.client_id}")
+ await self.close()
+
+ async def _handle_error(self, error):
+ """处理连接错误"""
+ logger.error(f"MQTT连接错误: {self.client_id}, error: {error}")
+ await self.close()
+
+ async def _keep_alive_check(self):
+ """心跳检查任务"""
+ try:
+ while self.is_connected_flag and not self._closed:
+ await asyncio.sleep(self.keep_alive_interval / 1000 / 2) # 检查间隔为心跳间隔的一半
+
+ current_time = time.time()
+ if current_time - self.last_activity > self.keep_alive_interval / 1000 * 1.5:
+ logger.info(f"MQTT客户端心跳超时: {self.client_id}")
+ await self.close()
+ break
+
+ except asyncio.CancelledError:
+ pass
+ except Exception as e:
+ logger.error(f"心跳检查任务出错: {e}")
+
+ def set_message_callback(self, callback: Callable[[str, str], None]):
+ """设置消息接收回调"""
+ self.message_callback = callback
+
+ async def send_message(self, topic: str, payload: str):
+ """发送MQTT消息"""
+ if self._closed or not self.is_connected_flag:
+ return
+
+ try:
+ await self.protocol.send_publish(topic, payload, qos=0)
+ logger.debug(f"发送MQTT消息: topic={topic}, payload={payload}")
+
+ except Exception as e:
+ logger.error(f"发送MQTT消息失败: {e}")
+
+ def is_connected(self) -> bool:
+ """检查连接状态"""
+ return self.is_connected_flag and not self._closed
+
+ async def close(self):
+ """关闭连接"""
+ if self._closed:
+ return
+
+ self._closed = True
+ self.is_connected_flag = False
+
+ # 取消心跳检查任务
+ if self.keep_alive_task and not self.keep_alive_task.done():
+ self.keep_alive_task.cancel()
+ try:
+ await self.keep_alive_task
+ except asyncio.CancelledError:
+ pass
+
+ # 通知服务器连接关闭
+ try:
+ await self.mqtt_server.on_client_disconnected(self)
+ except Exception as e:
+ logger.error(f"通知服务器连接关闭失败: {e}")
+
+ # 关闭协议处理器
+ try:
+ await self.protocol.close()
+ except Exception as e:
+ logger.error(f"关闭MQTT协议处理器失败: {e}")
+
+ logger.info(f"MQTT连接已关闭: {self.client_id}")
diff --git a/main/xiaozhi-server/core/protocols/mqtt_protocol.py b/main/xiaozhi-server/core/protocols/mqtt_protocol.py
new file mode 100644
index 00000000..f04e83b8
--- /dev/null
+++ b/main/xiaozhi-server/core/protocols/mqtt_protocol.py
@@ -0,0 +1,432 @@
+import asyncio
+from typing import Dict, Any, Callable
+from config.logger import setup_logging
+
+logger = setup_logging()
+
+
+# MQTT 固定头部的类型
+class PacketType:
+ CONNECT = 1
+ CONNACK = 2
+ PUBLISH = 3
+ SUBSCRIBE = 8
+ SUBACK = 9
+ PINGREQ = 12
+ PINGRESP = 13
+ DISCONNECT = 14
+
+
+class MQTTProtocol:
+ """
+ MQTT协议处理器:负责MQTT协议的解析和封装
+ """
+
+ def __init__(self, socket):
+ self.socket = socket
+ self.buffer = b''
+ self.event_handlers = {}
+ self.is_connected = False
+ self.keep_alive_interval = 0
+ self.last_activity = 0
+
+ # 启动消息处理任务
+ self._processing_task = asyncio.create_task(self._process_messages())
+
+ def on(self, event: str, handler: Callable):
+ """注册事件处理器"""
+ self.event_handlers[event] = handler
+
+ def emit(self, event: str, *args, **kwargs):
+ """触发事件"""
+ handler = self.event_handlers.get(event)
+ if handler:
+ if asyncio.iscoroutinefunction(handler):
+ asyncio.create_task(handler(*args, **kwargs))
+ else:
+ handler(*args, **kwargs)
+
+ async def _process_messages(self):
+ """处理消息的主循环"""
+ try:
+ while True:
+ # 从socket读取数据
+ data = await self._read_socket()
+ if not data:
+ break
+
+ # 添加到缓冲区
+ self.buffer += data
+
+ # 处理缓冲区中的消息
+ await self._process_buffer()
+
+ except asyncio.CancelledError:
+ pass
+ except Exception as e:
+ logger.error(f"MQTT消息处理循环出错: {e}")
+ self.emit('error', e)
+ finally:
+ self.emit('close')
+
+ async def _read_socket(self) -> bytes:
+ """从socket读取数据"""
+ try:
+ # 使用asyncio的socket读取
+ loop = asyncio.get_event_loop()
+ data = await loop.sock_recv(self.socket, 4096)
+ return data
+ except Exception as e:
+ logger.error(f"读取socket数据失败: {e}")
+ return b''
+
+ async def _process_buffer(self):
+ """处理缓冲区中的消息"""
+ while len(self.buffer) >= 2: # 至少需要2字节开始解析
+ try:
+ # 解析消息
+ message_length, message = self._parse_message()
+ if message_length == 0:
+ break # 消息不完整,等待更多数据
+
+ # 从缓冲区移除已处理的消息
+ self.buffer = self.buffer[message_length:]
+
+ # 处理消息
+ await self._handle_message(message)
+
+ except Exception as e:
+ logger.error(f"处理MQTT消息失败: {e}")
+ self.emit('protocolError', e)
+ break
+
+ def _parse_message(self) -> tuple[int, Dict[str, Any]]:
+ """解析MQTT消息"""
+ if len(self.buffer) < 2:
+ return 0, {}
+
+ # 获取消息类型
+ first_byte = self.buffer[0]
+ packet_type = (first_byte >> 4)
+
+ # 解析剩余长度
+ remaining_length, bytes_read = self._decode_remaining_length()
+ if remaining_length == -1:
+ return 0, {} # 长度解析失败,等待更多数据
+
+ # 计算完整消息长度
+ total_length = 1 + bytes_read + remaining_length
+
+ if len(self.buffer) < total_length:
+ return 0, {} # 消息不完整
+
+ # 提取消息数据
+ message_data = self.buffer[:total_length]
+
+ # 根据消息类型解析
+ if packet_type == PacketType.CONNECT:
+ message = self._parse_connect(message_data)
+ elif packet_type == PacketType.PUBLISH:
+ message = self._parse_publish(message_data)
+ elif packet_type == PacketType.SUBSCRIBE:
+ message = self._parse_subscribe(message_data)
+ elif packet_type == PacketType.PINGREQ:
+ message = {'type': 'pingreq'}
+ elif packet_type == PacketType.DISCONNECT:
+ message = {'type': 'disconnect'}
+ else:
+ logger.warning(f"未处理的MQTT消息类型: {packet_type}")
+ message = {'type': 'unknown', 'packet_type': packet_type}
+
+ return total_length, message
+
+ def _decode_remaining_length(self) -> tuple[int, int]:
+ """解码剩余长度字段"""
+ multiplier = 1
+ value = 0
+ bytes_read = 0
+
+ while bytes_read < 4 and bytes_read + 1 < len(self.buffer):
+ digit = self.buffer[bytes_read + 1]
+ bytes_read += 1
+
+ value += (digit & 127) * multiplier
+ multiplier *= 128
+
+ if (digit & 128) == 0:
+ break
+ else:
+ if bytes_read >= 4:
+ return -1, 0 # 长度字段过长
+ return -1, 0 # 数据不完整
+
+ return value, bytes_read
+
+ def _encode_remaining_length(self, length: int) -> bytes:
+ """编码剩余长度字段"""
+ result = bytearray()
+
+ while True:
+ digit = length % 128
+ length = length // 128
+
+ if length > 0:
+ digit |= 0x80
+
+ result.append(digit)
+
+ if length == 0:
+ break
+
+ return bytes(result)
+
+ def _parse_connect(self, message_data: bytes) -> Dict[str, Any]:
+ """解析CONNECT消息"""
+ try:
+ # 跳过固定头部和剩余长度
+ _, bytes_read = self._decode_remaining_length()
+ pos = 1 + bytes_read
+
+ # 协议名长度
+ protocol_length = int.from_bytes(message_data[pos:pos+2], 'big')
+ pos += 2
+
+ # 协议名
+ protocol = message_data[pos:pos+protocol_length].decode('utf-8')
+ pos += protocol_length
+
+ # 协议级别
+ protocol_level = message_data[pos]
+ pos += 1
+
+ # 连接标志
+ connect_flags = message_data[pos]
+ has_username = (connect_flags & 0x80) != 0
+ has_password = (connect_flags & 0x40) != 0
+ pos += 1
+
+ # 保持连接时间
+ keep_alive = int.from_bytes(message_data[pos:pos+2], 'big')
+ pos += 2
+
+ # 客户端ID
+ client_id_length = int.from_bytes(message_data[pos:pos+2], 'big')
+ pos += 2
+ client_id = message_data[pos:pos+client_id_length].decode('utf-8')
+ pos += client_id_length
+
+ # 用户名(如果存在)
+ username = ''
+ if has_username:
+ username_length = int.from_bytes(message_data[pos:pos+2], 'big')
+ pos += 2
+ username = message_data[pos:pos+username_length].decode('utf-8')
+ pos += username_length
+
+ # 密码(如果存在)
+ password = ''
+ if has_password:
+ password_length = int.from_bytes(message_data[pos:pos+2], 'big')
+ pos += 2
+ password = message_data[pos:pos+password_length].decode('utf-8')
+ pos += password_length
+
+ return {
+ 'type': 'connect',
+ 'protocol': protocol,
+ 'protocolLevel': protocol_level,
+ 'clientId': client_id,
+ 'keepAlive': keep_alive,
+ 'username': username,
+ 'password': password
+ }
+
+ except Exception as e:
+ logger.error(f"解析CONNECT消息失败: {e}")
+ raise
+
+ def _parse_publish(self, message_data: bytes) -> Dict[str, Any]:
+ """解析PUBLISH消息"""
+ try:
+ # 获取QoS等标志
+ first_byte = message_data[0]
+ qos = (first_byte & 0x06) >> 1
+ dup = (first_byte & 0x08) != 0
+ retain = (first_byte & 0x01) != 0
+
+ # 跳过固定头部和剩余长度
+ _, bytes_read = self._decode_remaining_length()
+ pos = 1 + bytes_read
+
+ # 主题长度
+ topic_length = int.from_bytes(message_data[pos:pos+2], 'big')
+ pos += 2
+
+ # 主题
+ topic = message_data[pos:pos+topic_length].decode('utf-8')
+ pos += topic_length
+
+ # 消息ID(QoS > 0时存在)
+ packet_id = None
+ if qos > 0:
+ packet_id = int.from_bytes(message_data[pos:pos+2], 'big')
+ pos += 2
+
+ # 有效载荷
+ payload = message_data[pos:].decode('utf-8')
+
+ return {
+ 'type': 'publish',
+ 'topic': topic,
+ 'payload': payload,
+ 'qos': qos,
+ 'dup': dup,
+ 'retain': retain,
+ 'packetId': packet_id
+ }
+
+ except Exception as e:
+ logger.error(f"解析PUBLISH消息失败: {e}")
+ raise
+
+ def _parse_subscribe(self, message_data: bytes) -> Dict[str, Any]:
+ """解析SUBSCRIBE消息"""
+ try:
+ # 跳过固定头部和剩余长度
+ _, bytes_read = self._decode_remaining_length()
+ pos = 1 + bytes_read
+
+ # 消息ID
+ packet_id = int.from_bytes(message_data[pos:pos+2], 'big')
+ pos += 2
+
+ # 主题长度
+ topic_length = int.from_bytes(message_data[pos:pos+2], 'big')
+ pos += 2
+
+ # 主题
+ topic = message_data[pos:pos+topic_length].decode('utf-8')
+ pos += topic_length
+
+ # QoS
+ qos = message_data[pos]
+
+ return {
+ 'type': 'subscribe',
+ 'packetId': packet_id,
+ 'topic': topic,
+ 'qos': qos
+ }
+
+ except Exception as e:
+ logger.error(f"解析SUBSCRIBE消息失败: {e}")
+ raise
+
+ async def _handle_message(self, message: Dict[str, Any]):
+ """处理解析后的消息"""
+ message_type = message.get('type')
+
+ if message_type == 'connect':
+ self.keep_alive_interval = message.get('keepAlive', 0)
+ self.is_connected = True
+ self.emit('connect', message)
+ elif message_type == 'publish':
+ self.emit('publish', message)
+ elif message_type == 'subscribe':
+ self.emit('subscribe', message)
+ elif message_type == 'pingreq':
+ await self.send_pingresp()
+ elif message_type == 'disconnect':
+ self.emit('disconnect')
+ else:
+ logger.warning(f"未处理的消息类型: {message_type}")
+
+ async def send_connack(self, return_code: int = 0, session_present: bool = False):
+ """发送CONNACK消息"""
+ packet = bytearray([
+ PacketType.CONNACK << 4, # 固定头部
+ 2, # 剩余长度
+ 1 if session_present else 0, # 连接确认标志
+ return_code # 返回码
+ ])
+
+ await self._send_packet(packet)
+
+ async def send_publish(self, topic: str, payload: str, qos: int = 0,
+ dup: bool = False, retain: bool = False, packet_id: int = None):
+ """发送PUBLISH消息"""
+ # 构造固定头部
+ first_byte = PacketType.PUBLISH << 4
+ if dup:
+ first_byte |= 0x08
+ if qos > 0:
+ first_byte |= (qos << 1)
+ if retain:
+ first_byte |= 0x01
+
+ # 构造可变头部和载荷
+ topic_bytes = topic.encode('utf-8')
+ payload_bytes = payload.encode('utf-8')
+
+ variable_header = bytearray()
+ variable_header.extend(len(topic_bytes).to_bytes(2, 'big'))
+ variable_header.extend(topic_bytes)
+
+ if qos > 0 and packet_id is not None:
+ variable_header.extend(packet_id.to_bytes(2, 'big'))
+
+ # 计算剩余长度
+ remaining_length = len(variable_header) + len(payload_bytes)
+ remaining_length_bytes = self._encode_remaining_length(remaining_length)
+
+ # 构造完整消息
+ packet = bytearray([first_byte])
+ packet.extend(remaining_length_bytes)
+ packet.extend(variable_header)
+ packet.extend(payload_bytes)
+
+ await self._send_packet(packet)
+
+ async def send_suback(self, packet_id: int, return_code: int = 0):
+ """发送SUBACK消息"""
+ packet = bytearray([
+ PacketType.SUBACK << 4, # 固定头部
+ 3, # 剩余长度
+ packet_id >> 8, # 消息ID高字节
+ packet_id & 0xFF, # 消息ID低字节
+ return_code # 返回码
+ ])
+
+ await self._send_packet(packet)
+
+ async def send_pingresp(self):
+ """发送PINGRESP消息"""
+ packet = bytearray([
+ PacketType.PINGRESP << 4, # 固定头部
+ 0 # 剩余长度
+ ])
+
+ await self._send_packet(packet)
+
+ async def _send_packet(self, packet: bytearray):
+ """发送数据包"""
+ try:
+ loop = asyncio.get_event_loop()
+ await loop.sock_sendall(self.socket, bytes(packet))
+ except Exception as e:
+ logger.error(f"发送MQTT数据包失败: {e}")
+ raise
+
+ async def close(self):
+ """关闭协议处理器"""
+ if hasattr(self, '_processing_task') and not self._processing_task.done():
+ self._processing_task.cancel()
+ try:
+ await self._processing_task
+ except asyncio.CancelledError:
+ pass
+
+ try:
+ self.socket.close()
+ except Exception as e:
+ logger.error(f"关闭socket失败: {e}")
diff --git a/main/xiaozhi-server/core/providers/asr/base.py b/main/xiaozhi-server/core/providers/asr/base.py
index 250a25f2..54bf643f 100644
--- a/main/xiaozhi-server/core/providers/asr/base.py
+++ b/main/xiaozhi-server/core/providers/asr/base.py
@@ -13,15 +13,80 @@ import concurrent.futures
from abc import ABC, abstractmethod
from config.logger import setup_logging
from typing import Optional, Tuple, List
-from core.handle.receiveAudioHandle import startToChat
-from core.handle.reportHandle import enqueue_asr_report
+# from core.handle.receiveAudioHandle import startToChat # 旧的handler
+# from core.handle.reportHandle import enqueue_asr_report # 旧的handler
+# from core.handle.receiveAudioHandle import handleAudioMessage # 旧的handler
+# 使用新的processor替代
from core.utils.util import remove_punctuation_and_length
-from core.handle.receiveAudioHandle import handleAudioMessage
TAG = __name__
logger = setup_logging()
+async def handleAudioMessage(conn, message):
+ """兼容函数:使用新的processor处理音频消息"""
+ try:
+ # 获取transport接口
+ transport = getattr(conn, 'transport', None)
+ if not transport:
+ logger.error("SessionContext中没有transport接口")
+ return
+
+ # 使用AudioReceiveProcessor处理音频消息
+ from core.processors.audio_receive_processor import AudioReceiveProcessor
+ processor = AudioReceiveProcessor()
+
+ # 处理音频消息
+ await processor.handle_audio_message(conn, transport, message)
+
+ except Exception as e:
+ logger.error(f"处理音频消息失败: {e}")
+ import traceback
+ traceback.print_exc()
+
+
+async def startToChat(conn, text):
+ """兼容函数:使用新的processor开始聊天"""
+ try:
+ # 获取transport接口
+ transport = getattr(conn, 'transport', None)
+ if not transport:
+ logger.error("SessionContext中没有transport接口")
+ return
+
+ # 使用ChatProcessor处理聊天
+ from core.processors.chat_processor import ChatProcessor
+ processor = ChatProcessor()
+
+ # 开始聊天
+ await processor.handle_chat(conn, transport, text)
+
+ except Exception as e:
+ logger.error(f"开始聊天失败: {e}")
+ import traceback
+ traceback.print_exc()
+
+
+def enqueue_asr_report(conn, text, audio_data):
+ """兼容函数:使用新的processor处理ASR报告"""
+ try:
+ # 获取transport接口
+ transport = getattr(conn, 'transport', None)
+ if not transport:
+ logger.error("SessionContext中没有transport接口")
+ return
+
+ # 使用ReportProcessor处理报告
+ from core.processors.report_processor import ReportProcessor
+ processor = ReportProcessor()
+
+ # 处理ASR报告
+ processor.enqueue_asr_report(conn, text, audio_data)
+
+ except Exception as e:
+ logger.error(f"ASR报告处理失败: {e}")
+
+
class ASRProviderBase(ABC):
def __init__(self):
pass
diff --git a/main/xiaozhi-server/core/providers/tools/device_iot/iot_executor.py b/main/xiaozhi-server/core/providers/tools/device_iot/iot_executor.py
index 7be9e112..6e5c998a 100644
--- a/main/xiaozhi-server/core/providers/tools/device_iot/iot_executor.py
+++ b/main/xiaozhi-server/core/providers/tools/device_iot/iot_executor.py
@@ -127,7 +127,16 @@ class DeviceIoTExecutor(ToolExecutor):
send_message = json.dumps(
{"type": "iot", "commands": [command]}
)
- await self.conn.websocket.send(send_message)
+
+ # 使用transport接口发送消息
+ if hasattr(self.conn, 'transport') and self.conn.transport:
+ await self.conn.transport.send(send_message)
+ elif hasattr(self.conn, 'websocket') and self.conn.websocket:
+ # 兼容旧版本
+ logger.warning("未找到SessionContext的传输层接口, 回退使用旧版conn.websocket发送消息")
+ await self.conn.websocket.send(send_message)
+ else:
+ raise AttributeError("无法找到可用的传输层接口")
return
raise Exception(f"未找到设备{device_name}的方法{method_name}")
diff --git a/main/xiaozhi-server/core/providers/tools/device_mcp/mcp_handler.py b/main/xiaozhi-server/core/providers/tools/device_mcp/mcp_handler.py
index c1b4d6c8..ad4e7aac 100644
--- a/main/xiaozhi-server/core/providers/tools/device_mcp/mcp_handler.py
+++ b/main/xiaozhi-server/core/providers/tools/device_mcp/mcp_handler.py
@@ -96,7 +96,7 @@ class MCPClient:
self.call_results.pop(id)
-async def send_mcp_message(conn, payload: dict):
+async def send_mcp_message(conn, payload: dict, transport=None):
"""Helper to send MCP messages, encapsulating common logic."""
if not conn.features.get("mcp"):
logger.bind(tag=TAG).warning("客户端不支持MCP,无法发送MCP消息")
@@ -105,13 +105,23 @@ async def send_mcp_message(conn, payload: dict):
message = json.dumps({"type": "mcp", "payload": payload})
try:
- await conn.websocket.send(message)
+ # 优先使用传入的transport,否则尝试从conn获取
+ if transport:
+ await transport.send(message)
+ elif hasattr(conn, 'websocket'):
+ # 兼容旧版本
+ await conn.websocket.send(message)
+ elif hasattr(conn, 'transport'):
+ # 新架构
+ await conn.transport.send(message)
+ else:
+ raise AttributeError("无法找到可用的传输层接口")
logger.bind(tag=TAG).info(f"成功发送MCP消息: {message}")
except Exception as e:
logger.bind(tag=TAG).error(f"发送MCP消息失败: {e}")
-async def handle_mcp_message(conn, mcp_client: MCPClient, payload: dict):
+async def handle_mcp_message(conn, mcp_client: MCPClient, payload: dict, transport=None):
"""处理MCP消息,包括初始化、工具列表和工具调用响应等"""
logger.bind(tag=TAG).info(f"处理MCP消息: {str(payload)[:100]}")
@@ -196,7 +206,7 @@ async def handle_mcp_message(conn, mcp_client: MCPClient, payload: dict):
next_cursor = result.get("nextCursor", "")
if next_cursor:
logger.bind(tag=TAG).info(f"有更多工具,nextCursor: {next_cursor}")
- await send_mcp_tools_list_continue_request(conn, next_cursor)
+ await send_mcp_tools_list_continue_request(conn, next_cursor, transport)
else:
await mcp_client.set_ready(True)
logger.bind(tag=TAG).info("所有工具已获取,MCP客户端准备就绪")
@@ -224,7 +234,7 @@ async def handle_mcp_message(conn, mcp_client: MCPClient, payload: dict):
)
-async def send_mcp_initialize_message(conn):
+async def send_mcp_initialize_message(conn, transport=None):
"""发送MCP初始化消息"""
vision_url = get_vision_url(conn.config)
@@ -256,10 +266,10 @@ async def send_mcp_initialize_message(conn):
},
}
logger.bind(tag=TAG).info("发送MCP初始化消息")
- await send_mcp_message(conn, payload)
+ await send_mcp_message(conn, payload, transport)
-async def send_mcp_tools_list_request(conn):
+async def send_mcp_tools_list_request(conn, transport=None):
"""发送MCP工具列表请求"""
payload = {
"jsonrpc": "2.0",
@@ -267,10 +277,10 @@ async def send_mcp_tools_list_request(conn):
"method": "tools/list",
}
logger.bind(tag=TAG).debug("发送MCP工具列表请求")
- await send_mcp_message(conn, payload)
+ await send_mcp_message(conn, payload, transport)
-async def send_mcp_tools_list_continue_request(conn, cursor: str):
+async def send_mcp_tools_list_continue_request(conn, cursor: str, transport=None):
"""发送带有cursor的MCP工具列表请求"""
payload = {
"jsonrpc": "2.0",
@@ -279,7 +289,7 @@ async def send_mcp_tools_list_continue_request(conn, cursor: str):
"params": {"cursor": cursor},
}
logger.bind(tag=TAG).info(f"发送带cursor的MCP工具列表请求: {cursor}")
- await send_mcp_message(conn, payload)
+ await send_mcp_message(conn, payload, transport)
async def call_mcp_tool(
diff --git a/main/xiaozhi-server/core/providers/tts/base.py b/main/xiaozhi-server/core/providers/tts/base.py
index 04a7fa36..0c47cb2c 100644
--- a/main/xiaozhi-server/core/providers/tts/base.py
+++ b/main/xiaozhi-server/core/providers/tts/base.py
@@ -14,8 +14,6 @@ from abc import ABC, abstractmethod
from config.logger import setup_logging
from core.utils.tts import MarkdownCleaner
from core.utils.output_counter import add_device_output
-from core.handle.reportHandle import enqueue_tts_report
-from core.handle.sendAudioHandle import sendAudioMessage
from core.utils.util import audio_bytes_to_data_stream, audio_to_data_stream
from core.providers.tts.dto.dto import (
TTSMessageDTO,
@@ -28,6 +26,72 @@ TAG = __name__
logger = setup_logging()
+async def sendAudioMessage(conn, sentenceType, audios, text):
+ """兼容函数:使用新的processor发送音频消息"""
+ try:
+ # 获取transport接口
+ transport = getattr(conn, 'transport', None)
+ if not transport:
+ logger.error("SessionContext中没有transport接口")
+ return
+
+ # 使用AudioSendProcessor发送音频
+ from core.processors.audio_send_processor import AudioSendProcessor
+ processor = AudioSendProcessor()
+
+ # 处理TTS开始消息
+ if conn.tts.tts_audio_first_sentence:
+ logger.info(f"发送第一段语音: {text}")
+ conn.tts.tts_audio_first_sentence = False
+ await processor.send_tts_message(conn, transport, "start", None)
+
+ if sentenceType == SentenceType.FIRST:
+ await processor.send_tts_message(conn, transport, "sentence_start", text)
+
+ await processor.send_audio(conn, transport, audios)
+
+ # 发送句子开始消息
+ if sentenceType is not SentenceType.MIDDLE:
+ logger.info(f"发送音频消息: {sentenceType}, {text}")
+
+ # 发送结束消息(如果是最后一个文本)
+ if conn.llm_finish_task and sentenceType == SentenceType.LAST:
+ await processor.send_tts_message(conn, transport, "stop", None)
+ conn.client_is_speaking = False
+ if conn.close_after_chat:
+ if hasattr(transport, 'close'):
+ await transport.close()
+
+ except Exception as e:
+ logger.error(f"发送音频消息失败: {e}")
+ import traceback
+ traceback.print_exc()
+
+
+def enqueue_tts_report(conn, audio_data, text):
+ """兼容函数:使用新的processor处理TTS报告"""
+ try:
+ # 获取transport接口
+ transport = getattr(conn, 'transport', None)
+ if not transport:
+ logger.error("SessionContext中没有transport接口")
+ return
+
+ # 使用ReportProcessor处理报告
+ from core.processors.report_processor import ReportProcessor
+ processor = ReportProcessor()
+
+ # 异步执行报告
+ if hasattr(conn, 'loop') and conn.loop:
+ # 直接调用同步方法
+ processor.enqueue_tts_report(conn, text, audio_data)
+ else:
+ logger.warning("SessionContext中没有事件循环,跳过TTS报告")
+
+ except Exception as e:
+ logger.error(f"TTS报告处理失败: {e}")
+
+
class TTSProviderBase(ABC):
def __init__(self, config, delete_audio_file):
self.interface_type = InterfaceType.NON_STREAM
diff --git a/main/xiaozhi-server/core/servers/mqtt_server.py b/main/xiaozhi-server/core/servers/mqtt_server.py
new file mode 100644
index 00000000..d5b901e4
--- /dev/null
+++ b/main/xiaozhi-server/core/servers/mqtt_server.py
@@ -0,0 +1,313 @@
+import asyncio
+import socket
+import time
+from typing import Dict, Any, Set
+from config.logger import setup_logging
+from core.protocols.mqtt_connection import MQTTConnection
+from core.transport.mqtt_transport import MQTTTransport, UDPAudioHandler
+from core.services.connection_service import ConnectionService
+
+logger = setup_logging()
+
+
+class MQTTServer:
+ """
+ 原生MQTT服务器:直接处理MQTT协议连接
+ 集成到xiaozhi-server架构中
+ """
+
+ def __init__(self, config: Dict[str, Any]):
+ self.config = config
+ self.logger = setup_logging()
+
+ # 服务器配置
+ server_config = config.get('mqtt_server', {})
+ self.mqtt_port = server_config.get('port', 1883)
+ self.udp_port = server_config.get('udp_port', self.mqtt_port)
+ self.host = server_config.get('host', '0.0.0.0')
+ self.public_ip = server_config.get('public_ip', 'localhost')
+
+ # 连接管理
+ self.connections: Dict[int, MQTTConnection] = {}
+ self.udp_handlers: Dict[int, UDPAudioHandler] = {}
+ self.connection_id_counter = 0
+
+ # 服务器实例
+ self.mqtt_server = None
+ self.udp_server = None
+
+ # 连接服务
+ self.connection_service = ConnectionService(config)
+
+ # 活跃连接管理
+ self.active_transports: Set[MQTTTransport] = set()
+
+ # 心跳检查
+ self.heartbeat_task = None
+ self.heartbeat_interval = 30 # 30秒检查一次
+
+ async def start(self):
+ """启动MQTT服务器"""
+ try:
+ # 启动MQTT TCP服务器
+ await self._start_mqtt_server()
+
+ # 启动UDP服务器
+ await self._start_udp_server()
+
+ # 启动心跳检查
+ self.heartbeat_task = asyncio.create_task(self._heartbeat_check())
+
+ logger.info(f"MQTT服务器启动成功: {self.host}:{self.mqtt_port}")
+ logger.info(f"UDP服务器启动成功: {self.host}:{self.udp_port}")
+
+ except Exception as e:
+ logger.error(f"启动MQTT服务器失败: {e}")
+ raise
+
+ async def _start_mqtt_server(self):
+ """启动MQTT TCP服务器"""
+ self.mqtt_server = await asyncio.start_server(
+ self._handle_mqtt_connection,
+ self.host,
+ self.mqtt_port
+ )
+
+ async def _start_udp_server(self):
+ """启动UDP服务器"""
+ loop = asyncio.get_event_loop()
+
+ # 创建UDP socket
+ sock = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
+ sock.bind((self.host, self.udp_port))
+ sock.setblocking(False)
+
+ # 创建UDP协议处理器
+ transport, protocol = await loop.create_datagram_endpoint(
+ lambda: UDPProtocol(self),
+ sock=sock
+ )
+
+ self.udp_server = (transport, protocol)
+
+ async def _handle_mqtt_connection(self, reader, writer):
+ """处理新的MQTT连接"""
+ connection_id = self._generate_connection_id()
+
+ try:
+ # 获取客户端地址
+ client_addr = writer.get_extra_info('peername')
+ logger.info(f"新MQTT连接: {client_addr}, connection_id: {connection_id}")
+
+ # 创建MQTT连接处理器
+ mqtt_connection = MQTTConnection(
+ writer.get_extra_info('socket'),
+ connection_id,
+ self
+ )
+
+ # 创建UDP音频处理器
+ udp_handler = UDPAudioHandler(
+ connection_id,
+ self,
+ {} # 加密配置将在hello消息中设置
+ )
+
+ # 创建MQTT传输层
+ transport = MQTTTransport(mqtt_connection, udp_handler)
+
+ # 注册连接
+ self.connections[connection_id] = mqtt_connection
+ self.udp_handlers[connection_id] = udp_handler
+ self.active_transports.add(transport)
+
+ # 提取连接头信息
+ headers = {
+ 'x-real-ip': client_addr[0] if client_addr else 'unknown',
+ 'connection-type': 'mqtt'
+ }
+
+ try:
+ # 使用ConnectionService处理连接
+ await self.connection_service.handle_connection(transport, headers)
+
+ except Exception as e:
+ logger.error(f"ConnectionService处理MQTT连接失败: {e}")
+
+ except Exception as e:
+ logger.error(f"处理MQTT连接失败: {e}")
+ finally:
+ # 清理连接
+ await self._cleanup_connection(connection_id)
+
+ async def _cleanup_connection(self, connection_id: int):
+ """清理连接资源"""
+ try:
+ # 移除连接
+ if connection_id in self.connections:
+ connection = self.connections.pop(connection_id)
+ await connection.close()
+
+ # 移除UDP处理器
+ if connection_id in self.udp_handlers:
+ udp_handler = self.udp_handlers.pop(connection_id)
+ await udp_handler.close()
+
+ # 移除传输层(通过连接ID查找)
+ transports_to_remove = []
+ for transport in self.active_transports:
+ if hasattr(transport, '_mqtt_connection') and \
+ transport._mqtt_connection.connection_id == connection_id:
+ transports_to_remove.append(transport)
+
+ for transport in transports_to_remove:
+ self.active_transports.discard(transport)
+ await transport.close()
+
+ logger.info(f"MQTT连接清理完成: {connection_id}")
+
+ except Exception as e:
+ logger.error(f"清理MQTT连接失败: {e}")
+
+ def _generate_connection_id(self) -> int:
+ """生成连接ID"""
+ self.connection_id_counter += 1
+ return self.connection_id_counter
+
+ async def on_client_connected(self, mqtt_connection: MQTTConnection):
+ """客户端连接回调"""
+ logger.info(f"MQTT客户端已连接: {mqtt_connection.client_id}")
+
+ async def on_client_disconnected(self, mqtt_connection: MQTTConnection):
+ """客户端断开连接回调"""
+ logger.info(f"MQTT客户端已断开: {mqtt_connection.client_id}")
+
+ async def send_udp_message(self, data: bytes, remote_addr: tuple):
+ """发送UDP消息"""
+ if self.udp_server:
+ transport, protocol = self.udp_server
+ transport.sendto(data, remote_addr)
+
+ async def send_encrypted_audio(self, connection_id: int, audio_data: bytes,
+ timestamp: int, remote_addr: tuple, encryption_config: Dict[str, Any]):
+ """发送加密音频数据"""
+ try:
+ # 这里应该实现音频数据加密逻辑
+ # 暂时直接发送原始数据
+ header = self._generate_udp_header(connection_id, len(audio_data), timestamp, 0)
+ message = header + audio_data
+
+ await self.send_udp_message(message, remote_addr)
+
+ except Exception as e:
+ logger.error(f"发送加密音频失败: {e}")
+
+ def _generate_udp_header(self, connection_id: int, length: int, timestamp: int, sequence: int) -> bytes:
+ """生成UDP消息头"""
+ header = bytearray(16)
+ header[0] = 1 # type
+ header[2:4] = length.to_bytes(2, 'big') # payload length
+ header[4:8] = connection_id.to_bytes(4, 'big') # connection id
+ header[8:12] = timestamp.to_bytes(4, 'big') # timestamp
+ header[12:16] = sequence.to_bytes(4, 'big') # sequence
+ return bytes(header)
+
+ async def _heartbeat_check(self):
+ """心跳检查任务"""
+ try:
+ while True:
+ await asyncio.sleep(self.heartbeat_interval)
+
+ # 检查所有连接的状态
+ dead_connections = []
+ for connection_id, connection in self.connections.items():
+ if not connection.is_connected():
+ dead_connections.append(connection_id)
+
+ # 清理死连接
+ for connection_id in dead_connections:
+ logger.info(f"清理死连接: {connection_id}")
+ await self._cleanup_connection(connection_id)
+
+ # 记录活跃连接数
+ active_count = len(self.connections)
+ if active_count > 0:
+ logger.info(f"MQTT活跃连接数: {active_count}")
+
+ except asyncio.CancelledError:
+ pass
+ except Exception as e:
+ logger.error(f"心跳检查任务出错: {e}")
+
+ async def stop(self):
+ """停止MQTT服务器"""
+ logger.info("正在停止MQTT服务器...")
+
+ # 停止心跳检查
+ if self.heartbeat_task and not self.heartbeat_task.done():
+ self.heartbeat_task.cancel()
+ try:
+ await self.heartbeat_task
+ except asyncio.CancelledError:
+ pass
+
+ # 关闭所有连接
+ for connection_id in list(self.connections.keys()):
+ await self._cleanup_connection(connection_id)
+
+ # 关闭UDP服务器
+ if self.udp_server:
+ transport, protocol = self.udp_server
+ transport.close()
+
+ # 关闭MQTT服务器
+ if self.mqtt_server:
+ self.mqtt_server.close()
+ await self.mqtt_server.wait_closed()
+
+ logger.info("MQTT服务器已停止")
+
+ def get_server_status(self) -> Dict[str, Any]:
+ """获取服务器状态"""
+ return {
+ 'type': 'mqtt',
+ 'host': self.host,
+ 'mqtt_port': self.mqtt_port,
+ 'udp_port': self.udp_port,
+ 'active_connections': len(self.connections),
+ 'active_transports': len(self.active_transports)
+ }
+
+
+class UDPProtocol(asyncio.DatagramProtocol):
+ """UDP协议处理器"""
+
+ def __init__(self, mqtt_server: MQTTServer):
+ self.mqtt_server = mqtt_server
+ self.transport = None
+
+ def connection_made(self, transport):
+ self.transport = transport
+
+ def datagram_received(self, data: bytes, addr: tuple):
+ """接收UDP数据报"""
+ try:
+ # 解析UDP消息头
+ if len(data) < 16:
+ return
+
+ connection_id = int.from_bytes(data[4:8], 'big')
+ timestamp = int.from_bytes(data[8:12], 'big')
+ sequence = int.from_bytes(data[12:16], 'big')
+ payload = data[16:]
+
+ # 找到对应的UDP处理器
+ udp_handler = self.mqtt_server.udp_handlers.get(connection_id)
+ if udp_handler:
+ udp_handler.on_udp_message(payload, timestamp, addr)
+
+ except Exception as e:
+ logger.error(f"处理UDP数据报失败: {e}")
+
+ 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
new file mode 100644
index 00000000..57ffd4fd
--- /dev/null
+++ b/main/xiaozhi-server/core/servers/multi_protocol_server.py
@@ -0,0 +1,318 @@
+import asyncio
+from typing import Dict, Any, List, Optional
+from config.logger import setup_logging
+from core.websocket_server_new import NewWebSocketServer
+from core.servers.mqtt_server import MQTTServer
+
+logger = setup_logging()
+
+
+class MultiProtocolServer:
+ """
+ 多协议服务器管理器:统一管理WebSocket和MQTT服务器
+ 提供统一的启动、停止和状态监控接口
+ """
+
+ def __init__(self, config: Dict[str, Any]):
+ self.config = config
+ self.logger = setup_logging()
+
+ # 服务器实例
+ self.servers: Dict[str, Any] = {}
+ self.server_tasks: Dict[str, asyncio.Task] = {}
+
+ # 服务器状态
+ self.is_running = False
+ self.startup_complete = False
+
+ # 初始化服务器
+ self._initialize_servers()
+
+ def _initialize_servers(self):
+ """初始化所有协议服务器"""
+ try:
+ # 检查配置中启用的协议
+ enabled_protocols = self.config.get('enabled_protocols', ['websocket'])
+
+ # 初始化WebSocket服务器
+ if 'websocket' in enabled_protocols:
+ self.servers['websocket'] = NewWebSocketServer(self.config)
+ logger.info("WebSocket服务器已初始化")
+
+ # 初始化MQTT服务器
+ if 'mqtt' in enabled_protocols:
+ self.servers['mqtt'] = MQTTServer(self.config)
+ logger.info("MQTT服务器已初始化")
+
+ if not self.servers:
+ logger.warning("没有启用任何协议服务器")
+
+ except Exception as e:
+ logger.error(f"初始化服务器失败: {e}")
+ raise
+
+ async def start(self):
+ """启动所有服务器"""
+ if self.is_running:
+ logger.warning("服务器已经在运行中")
+ return
+
+ try:
+ logger.info("开始启动多协议服务器...")
+ self.is_running = True
+
+ # 启动所有服务器
+ for protocol, server in self.servers.items():
+ try:
+ logger.info(f"启动{protocol}服务器...")
+ task = asyncio.create_task(server.start())
+ self.server_tasks[protocol] = task
+
+ # 等待一小段时间确保服务器启动
+ await asyncio.sleep(0.1)
+
+ logger.info(f"{protocol}服务器启动成功")
+
+ except Exception as e:
+ logger.error(f"启动{protocol}服务器失败: {e}")
+ # 继续启动其他服务器
+ continue
+
+ self.startup_complete = True
+ logger.info("多协议服务器启动完成")
+
+ # 启动监控任务
+ asyncio.create_task(self._monitor_servers())
+
+ # 等待所有服务器任务
+ if self.server_tasks:
+ await asyncio.gather(*self.server_tasks.values(), return_exceptions=True)
+
+ except Exception as e:
+ logger.error(f"启动多协议服务器失败: {e}")
+ await self.stop()
+ raise
+
+ async def stop(self):
+ """停止所有服务器"""
+ if not self.is_running:
+ return
+
+ logger.info("开始停止多协议服务器...")
+ self.is_running = False
+
+ # 停止所有服务器
+ for protocol, server in self.servers.items():
+ try:
+ logger.info(f"停止{protocol}服务器...")
+ await server.stop()
+ logger.info(f"{protocol}服务器已停止")
+ except Exception as e:
+ logger.error(f"停止{protocol}服务器失败: {e}")
+
+ # 取消所有服务器任务
+ for protocol, task in self.server_tasks.items():
+ if not task.done():
+ task.cancel()
+ try:
+ await task
+ except asyncio.CancelledError:
+ pass
+
+ self.server_tasks.clear()
+ logger.info("多协议服务器已停止")
+
+ async def restart(self):
+ """重启所有服务器"""
+ logger.info("重启多协议服务器...")
+ await self.stop()
+ await asyncio.sleep(1) # 等待清理完成
+ await self.start()
+
+ async def restart_server(self, protocol: str):
+ """重启指定协议的服务器"""
+ if protocol not in self.servers:
+ logger.error(f"未找到协议服务器: {protocol}")
+ return False
+
+ try:
+ logger.info(f"重启{protocol}服务器...")
+
+ # 停止指定服务器
+ server = self.servers[protocol]
+ await server.stop()
+
+ # 取消任务
+ if protocol in self.server_tasks:
+ task = self.server_tasks[protocol]
+ if not task.done():
+ task.cancel()
+ try:
+ await task
+ except asyncio.CancelledError:
+ pass
+
+ # 重新启动
+ task = asyncio.create_task(server.start())
+ self.server_tasks[protocol] = task
+
+ logger.info(f"{protocol}服务器重启成功")
+ return True
+
+ except Exception as e:
+ logger.error(f"重启{protocol}服务器失败: {e}")
+ return False
+
+ async def update_config(self, new_config: Dict[str, Any]):
+ """更新配置"""
+ try:
+ logger.info("更新多协议服务器配置...")
+
+ # 检查配置变化
+ config_changed = self._check_config_changes(self.config, new_config)
+
+ # 更新配置
+ self.config = new_config
+
+ # 如果配置有重大变化,重新初始化服务器
+ if config_changed:
+ logger.info("配置有重大变化,重新初始化服务器...")
+ await self.stop()
+ self._initialize_servers()
+ if self.is_running:
+ await self.start()
+ else:
+ # 更新各个服务器的配置
+ for protocol, server in self.servers.items():
+ if hasattr(server, 'update_config'):
+ await server.update_config(new_config)
+
+ logger.info("配置更新完成")
+ return True
+
+ except Exception as e:
+ logger.error(f"更新配置失败: {e}")
+ return False
+
+ def _check_config_changes(self, old_config: Dict[str, Any], new_config: Dict[str, Any]) -> bool:
+ """检查配置是否有重大变化"""
+ # 检查启用的协议是否变化
+ old_protocols = set(old_config.get('enabled_protocols', ['websocket']))
+ new_protocols = set(new_config.get('enabled_protocols', ['websocket']))
+
+ if old_protocols != new_protocols:
+ logger.info(f"启用协议发生变化: {old_protocols} -> {new_protocols}")
+ return True
+
+ # 检查服务器端口配置
+ server_configs = ['server', 'mqtt_server']
+ for config_key in server_configs:
+ old_server_config = old_config.get(config_key, {})
+ new_server_config = new_config.get(config_key, {})
+
+ # 检查端口和主机配置
+ for key in ['port', 'host', 'ip']:
+ if old_server_config.get(key) != new_server_config.get(key):
+ logger.info(f"服务器配置{config_key}.{key}发生变化")
+ return True
+
+ return False
+
+ async def _monitor_servers(self):
+ """监控服务器状态"""
+ try:
+ while self.is_running:
+ await asyncio.sleep(30) # 每30秒检查一次
+
+ # 检查服务器任务状态
+ for protocol, task in self.server_tasks.items():
+ if task.done():
+ exception = task.exception()
+ if exception:
+ logger.error(f"{protocol}服务器异常退出: {exception}")
+ # 尝试重启服务器
+ await self.restart_server(protocol)
+
+ except asyncio.CancelledError:
+ pass
+ except Exception as e:
+ logger.error(f"服务器监控任务出错: {e}")
+
+ def get_server_status(self) -> Dict[str, Any]:
+ """获取所有服务器状态"""
+ status = {
+ 'is_running': self.is_running,
+ 'startup_complete': self.startup_complete,
+ 'enabled_protocols': list(self.servers.keys()),
+ 'servers': {}
+ }
+
+ # 获取各个服务器的状态
+ for protocol, server in self.servers.items():
+ try:
+ if hasattr(server, 'get_server_status'):
+ server_status = server.get_server_status()
+ else:
+ server_status = {'type': protocol, 'status': 'unknown'}
+
+ # 添加任务状态
+ task = self.server_tasks.get(protocol)
+ if task:
+ server_status['task_status'] = 'running' if not task.done() else 'stopped'
+ if task.done() and task.exception():
+ server_status['task_error'] = str(task.exception())
+
+ status['servers'][protocol] = server_status
+
+ except Exception as e:
+ status['servers'][protocol] = {
+ 'type': protocol,
+ 'status': 'error',
+ 'error': str(e)
+ }
+
+ return status
+
+ def get_active_connections_count(self) -> Dict[str, int]:
+ """获取各协议的活跃连接数"""
+ connections = {}
+
+ for protocol, server in self.servers.items():
+ try:
+ if hasattr(server, 'get_active_connections_count'):
+ connections[protocol] = server.get_active_connections_count()
+ elif hasattr(server, 'connections'):
+ connections[protocol] = len(server.connections)
+ else:
+ connections[protocol] = 0
+ except Exception as e:
+ logger.error(f"获取{protocol}连接数失败: {e}")
+ connections[protocol] = -1
+
+ return connections
+
+ async def broadcast_message(self, message: Dict[str, Any], protocol: Optional[str] = None):
+ """向所有连接广播消息"""
+ try:
+ if protocol:
+ # 向指定协议广播
+ if protocol in self.servers:
+ server = self.servers[protocol]
+ if hasattr(server, 'broadcast_message'):
+ await server.broadcast_message(message)
+ else:
+ # 向所有协议广播
+ for server in self.servers.values():
+ if hasattr(server, 'broadcast_message'):
+ await server.broadcast_message(message)
+
+ except Exception as e:
+ logger.error(f"广播消息失败: {e}")
+
+ def get_supported_protocols(self) -> List[str]:
+ """获取支持的协议列表"""
+ return ['websocket', 'mqtt']
+
+ 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
new file mode 100644
index 00000000..d1f4ea09
--- /dev/null
+++ b/main/xiaozhi-server/core/services/connection_service.py
@@ -0,0 +1,161 @@
+import asyncio
+from typing import Dict, Any
+from urllib.parse import parse_qs, urlparse
+
+from core.context.session_context import SessionContext
+from core.transport.transport_interface import TransportInterface
+from core.components.component_registry import ComponentRegistry
+from core.components.component_manager import ComponentType
+from core.pipeline.message_pipeline import MessagePipeline
+from core.processors.message_router import MessageRouter
+from config.logger import setup_logging
+
+logger = setup_logging()
+
+
+class ConnectionService:
+ """连接服务:统一管理连接生命周期,替代ConnectionHandler"""
+
+ def __init__(self, config: Dict[str, Any]):
+ self.config = config
+ self.logger = setup_logging()
+
+ # 创建统一消息路由器
+ self.message_router = MessageRouter()
+
+ # 创建消息处理管道
+ self.message_pipeline = MessagePipeline()
+ self._setup_pipeline()
+
+ def _setup_pipeline(self):
+ """设置消息处理管道"""
+ # 使用统一的MessageRouter替代单独的processor
+ self.message_pipeline.add_processor(self.message_router)
+
+ async def handle_connection(self, transport: TransportInterface, headers: Dict[str, str]):
+ """处理新连接"""
+ # 创建会话上下文
+ context = SessionContext()
+ context.config = self.config
+ context.headers = headers
+
+ # 设置transport接口
+ context.transport = transport
+
+ # 兼容性:设置websocket属性(如果transport是WebSocket)
+ if hasattr(transport, '_websocket'):
+ context.websocket = transport._websocket
+
+ # 从headers或URL参数中提取设备信息
+ await self._extract_device_info(context, headers)
+
+ # 创建组件管理器
+ component_manager = ComponentRegistry.create_component_manager(self.config)
+
+ # 启动超时检查任务
+ timeout_task = None
+
+ try:
+ logger.info(f"新连接建立: {context.device_id} from {context.client_ip}")
+
+ # 初始化必要的组件
+ await self._initialize_components(context, component_manager)
+
+ # 启动超时检查
+ timeout_task = asyncio.create_task(self._check_timeout(context, transport))
+
+ # 处理消息流
+ async for message in transport.receive():
+ try:
+ await self.message_pipeline.process_message(context, transport, message)
+ except Exception as e:
+ logger.error(f"处理消息时出错: {e}")
+ # 继续处理其他消息,不中断连接
+
+ except Exception as e:
+ logger.error(f"连接处理出错: {e}")
+ finally:
+ # 清理资源
+ if timeout_task and not timeout_task.done():
+ timeout_task.cancel()
+ try:
+ await timeout_task
+ except asyncio.CancelledError:
+ pass
+
+ # 清理组件
+ try:
+ await component_manager.cleanup_all()
+ except Exception as e:
+ logger.error(f"组件清理失败: {e}")
+
+ # 执行会话清理回调
+ try:
+ await context.run_cleanup()
+ except Exception as e:
+ logger.error(f"会话清理失败: {e}")
+
+ # 关闭传输层
+ try:
+ await transport.close()
+ except Exception as e:
+ logger.error(f"关闭传输层失败: {e}")
+
+ logger.info(f"连接已关闭: {context.device_id}")
+
+ async def _extract_device_info(self, context: SessionContext, headers: Dict[str, str]):
+ """从headers或URL参数中提取设备信息"""
+ device_id = headers.get("device-id")
+ client_id = headers.get("client-id")
+
+ # 如果headers中没有device-id,尝试从URL参数中获取
+ if not device_id:
+ # 这里需要从WebSocket请求中获取路径信息
+ # 暂时使用占位符,后续在WebSocketServer中传入
+ pass
+
+ context.device_id = device_id
+ context.client_ip = headers.get("x-real-ip") or headers.get("x-forwarded-for", "unknown")
+
+ if context.client_ip and "," in context.client_ip:
+ context.client_ip = context.client_ip.split(",")[0].strip()
+
+ async def _initialize_components(self, context: SessionContext, component_manager):
+ """初始化必要的组件"""
+ try:
+ # 设置组件管理器到上下文
+ context.component_manager = component_manager
+
+ # 根据配置确定需要初始化的组件
+ required_components = ComponentRegistry.get_required_components(self.config)
+
+ # 按需初始化组件
+ for component_type in required_components:
+ component = await component_manager.get_component(component_type, context)
+ if component:
+ logger.info(f"组件初始化成功: {component_type.value}")
+ else:
+ logger.warning(f"组件初始化失败: {component_type.value}")
+
+ except Exception as e:
+ logger.error(f"组件初始化出错: {e}")
+ raise
+
+ async def _check_timeout(self, context: SessionContext, transport: TransportInterface):
+ """定期检查连接超时"""
+ timeout_seconds = context.config.get("close_connection_no_voice_time", 120)
+ check_interval = min(30, timeout_seconds // 4) # 检查间隔为超时时间的1/4,最多30秒
+
+ try:
+ while transport.is_connected:
+ await asyncio.sleep(check_interval)
+
+ if context.is_timeout(timeout_seconds):
+ logger.info(f"连接超时,关闭连接: {context.session_id}")
+ await transport.close()
+ break
+
+ except asyncio.CancelledError:
+ pass
+ except Exception as e:
+ logger.error(f"超时检查出错: {e}")
diff --git a/main/xiaozhi-server/core/transport/mqtt_transport.py b/main/xiaozhi-server/core/transport/mqtt_transport.py
new file mode 100644
index 00000000..05eacb5c
--- /dev/null
+++ b/main/xiaozhi-server/core/transport/mqtt_transport.py
@@ -0,0 +1,245 @@
+import asyncio
+import json
+from typing import Any, AsyncGenerator, Dict, Optional
+from .transport_interface import TransportInterface
+from config.logger import setup_logging
+
+logger = setup_logging()
+
+
+class MQTTTransport(TransportInterface):
+ """
+ MQTT传输层实现:直接处理MQTT协议消息
+ 支持JSON消息和二进制音频数据传输
+ """
+
+ def __init__(self, mqtt_connection, udp_handler=None):
+ """
+ 初始化MQTT传输层
+
+ Args:
+ mqtt_connection: MQTT连接对象,包含协议处理器
+ udp_handler: UDP处理器,用于音频数据传输
+ """
+ self._mqtt_connection = mqtt_connection
+ self._udp_handler = udp_handler
+ self._message_queue = asyncio.Queue()
+ self._closed = False
+
+ # 设置MQTT连接的消息回调
+ self._setup_message_handlers()
+
+ def _setup_message_handlers(self):
+ """设置消息处理回调"""
+ # 设置MQTT消息接收回调
+ self._mqtt_connection.set_message_callback(self._on_mqtt_message)
+
+ # 设置UDP消息接收回调(如果有UDP处理器)
+ if self._udp_handler:
+ self._udp_handler.set_message_callback(self._on_udp_message)
+
+ def _on_mqtt_message(self, topic: str, payload: str):
+ """处理接收到的MQTT消息"""
+ try:
+ # 解析JSON消息
+ message_data = json.loads(payload)
+ message_data['_transport_type'] = 'mqtt'
+ message_data['_topic'] = topic
+
+ # 将消息放入队列
+ asyncio.create_task(self._message_queue.put(message_data))
+
+ except json.JSONDecodeError as e:
+ logger.error(f"MQTT消息JSON解析失败: {e}, payload: {payload}")
+ except Exception as e:
+ logger.error(f"处理MQTT消息失败: {e}")
+
+ def _on_udp_message(self, audio_data: bytes, timestamp: int):
+ """处理接收到的UDP音频消息"""
+ try:
+ # 构造音频消息格式
+ message_data = {
+ 'type': 'audio',
+ 'data': audio_data,
+ 'timestamp': timestamp,
+ '_transport_type': 'udp'
+ }
+
+ # 将消息放入队列
+ asyncio.create_task(self._message_queue.put(message_data))
+
+ except Exception as e:
+ logger.error(f"处理UDP音频消息失败: {e}")
+
+ async def send(self, data: Any) -> None:
+ """发送消息"""
+ if self._closed:
+ raise RuntimeError("Transport is closed")
+
+ try:
+ if isinstance(data, dict):
+ # 根据消息类型选择传输方式
+ if data.get('type') == 'audio' and self._udp_handler:
+ # 音频数据通过UDP发送
+ audio_data = data.get('data')
+ timestamp = data.get('timestamp', 0)
+ await self._udp_handler.send_audio(audio_data, timestamp)
+ else:
+ # JSON消息通过MQTT发送
+ topic = data.get('_topic', self._mqtt_connection.reply_topic)
+ payload = json.dumps(data)
+ await self._mqtt_connection.send_message(topic, payload)
+
+ elif isinstance(data, str):
+ # 字符串消息通过MQTT发送
+ await self._mqtt_connection.send_message(
+ self._mqtt_connection.reply_topic,
+ data
+ )
+
+ elif isinstance(data, bytes):
+ # 二进制数据通过UDP发送(如果有UDP处理器)
+ if self._udp_handler:
+ await self._udp_handler.send_audio(data, 0)
+ else:
+ logger.warning("尝试发送二进制数据但没有UDP处理器")
+
+ else:
+ # 其他类型转换为字符串通过MQTT发送
+ await self._mqtt_connection.send_message(
+ self._mqtt_connection.reply_topic,
+ str(data)
+ )
+
+ except Exception as e:
+ logger.error(f"MQTT传输发送消息失败: {e}")
+ raise
+
+ async def receive(self) -> AsyncGenerator[Any, None]:
+ """异步消息流"""
+ while not self._closed:
+ try:
+ # 等待消息,设置超时避免无限等待
+ message = await asyncio.wait_for(
+ self._message_queue.get(),
+ timeout=1.0
+ )
+ yield message
+
+ except asyncio.TimeoutError:
+ # 超时继续循环,检查连接状态
+ if not self.is_connected:
+ break
+ continue
+
+ except Exception as e:
+ logger.error(f"MQTT传输接收消息失败: {e}")
+ break
+
+ async def close(self) -> None:
+ """关闭传输层"""
+ if self._closed:
+ return
+
+ self._closed = True
+
+ try:
+ # 关闭MQTT连接
+ if self._mqtt_connection:
+ await self._mqtt_connection.close()
+
+ # 关闭UDP处理器
+ if self._udp_handler:
+ await self._udp_handler.close()
+
+ except Exception as e:
+ logger.error(f"关闭MQTT传输层失败: {e}")
+ raise RuntimeError("MQTT transport close failed")
+
+ @property
+ def is_connected(self) -> bool:
+ """检查连接状态"""
+ if self._closed:
+ return False
+
+ try:
+ # 检查MQTT连接状态
+ mqtt_connected = (
+ self._mqtt_connection and
+ self._mqtt_connection.is_connected()
+ )
+
+ return mqtt_connected
+
+ except Exception as e:
+ logger.error(f"检查MQTT连接状态失败: {e}")
+ return False
+
+ @property
+ def device_id(self) -> Optional[str]:
+ """获取设备ID"""
+ return getattr(self._mqtt_connection, 'device_id', None)
+
+ @property
+ def client_id(self) -> Optional[str]:
+ """获取客户端ID"""
+ return getattr(self._mqtt_connection, 'client_id', None)
+
+ @property
+ def session_id(self) -> Optional[str]:
+ """获取会话ID"""
+ return getattr(self._mqtt_connection, 'session_id', None)
+
+
+class UDPAudioHandler:
+ """
+ UDP音频处理器:处理加密音频数据传输
+ """
+
+ def __init__(self, connection_id: int, udp_server, encryption_config: Dict[str, Any]):
+ self.connection_id = connection_id
+ self.udp_server = udp_server
+ self.encryption_config = encryption_config
+ self.remote_address = None
+ self.message_callback = None
+ self._closed = False
+
+ def set_message_callback(self, callback):
+ """设置消息接收回调"""
+ self.message_callback = callback
+
+ async def send_audio(self, audio_data: bytes, timestamp: int):
+ """发送音频数据"""
+ if self._closed or not self.remote_address:
+ return
+
+ try:
+ # 使用UDP服务器发送加密音频数据
+ await self.udp_server.send_encrypted_audio(
+ self.connection_id,
+ audio_data,
+ timestamp,
+ self.remote_address,
+ self.encryption_config
+ )
+ except Exception as e:
+ logger.error(f"发送UDP音频数据失败: {e}")
+
+ def on_udp_message(self, audio_data: bytes, timestamp: int, remote_addr):
+ """处理接收到的UDP消息"""
+ if self._closed:
+ return
+
+ # 记录远程地址
+ if not self.remote_address:
+ self.remote_address = remote_addr
+
+ # 调用回调函数
+ if self.message_callback:
+ self.message_callback(audio_data, timestamp)
+
+ async def close(self):
+ """关闭UDP处理器"""
+ self._closed = True
+ self.message_callback = None
+ self.remote_address = None
diff --git a/main/xiaozhi-server/core/transport/transport_interface.py b/main/xiaozhi-server/core/transport/transport_interface.py
new file mode 100644
index 00000000..285f87f4
--- /dev/null
+++ b/main/xiaozhi-server/core/transport/transport_interface.py
@@ -0,0 +1,31 @@
+from abc import ABC, abstractmethod
+from typing import Any, AsyncGenerator
+
+
+class TransportInterface(ABC):
+ """
+ 传输层抽象接口。
+ """
+
+ @abstractmethod
+ async def send(self, data: Any) -> None:
+ """发送一条消息。"""
+ raise NotImplementedError
+
+ @abstractmethod
+ async def receive(self) -> AsyncGenerator[Any, None]:
+ """异步消息流。"""
+ yield # pragma: no cover
+
+ @abstractmethod
+ async def close(self) -> None:
+ """关闭底层连接。"""
+ raise NotImplementedError
+
+ @property
+ @abstractmethod
+ def is_connected(self) -> bool:
+ """连接是否存活。"""
+ raise NotImplementedError
+
+
diff --git a/main/xiaozhi-server/core/transport/websocket_transport.py b/main/xiaozhi-server/core/transport/websocket_transport.py
new file mode 100644
index 00000000..cbf704a5
--- /dev/null
+++ b/main/xiaozhi-server/core/transport/websocket_transport.py
@@ -0,0 +1,46 @@
+from typing import Any, AsyncGenerator
+from .transport_interface import TransportInterface
+
+
+class WebSocketTransport(TransportInterface):
+ """
+ WebSocket 传输实现:包装 websockets 库的协议对象,
+ 提供统一的 send/receive/close 接口。
+ """
+
+ def __init__(self, websocket):
+ self._ws = websocket
+
+ async def send(self, data: Any) -> None:
+ if isinstance(data, (str, bytes)):
+ await self._ws.send(data)
+ else:
+ await self._ws.send(str(data))
+
+ async def receive(self) -> AsyncGenerator[Any, None]:
+ async for message in self._ws:
+ yield message
+
+ async def close(self) -> None:
+ try:
+ if hasattr(self._ws, "closed") and not self._ws.closed:
+ await self._ws.close()
+ elif hasattr(self._ws, "state") and self._ws.state.name != "CLOSED":
+ await self._ws.close()
+ else:
+ await self._ws.close()
+ except Exception:
+ raise RuntimeError("WebSocket close failed")
+
+ @property
+ def is_connected(self) -> bool:
+ try:
+ if hasattr(self._ws, "closed"):
+ return not self._ws.closed
+ if hasattr(self._ws, "state"):
+ return getattr(self._ws.state, "name", "CLOSED") != "CLOSED"
+ except Exception:
+ raise RuntimeError("WebSocket connection check failed")
+ return False
+
+
diff --git a/main/xiaozhi-server/core/utils/textUtils.py b/main/xiaozhi-server/core/utils/textUtils.py
index 47677b99..21dac597 100644
--- a/main/xiaozhi-server/core/utils/textUtils.py
+++ b/main/xiaozhi-server/core/utils/textUtils.py
@@ -1,4 +1,5 @@
import json
+from config.logger import setup_logging
TAG = __name__
EMOJI_MAP = {
@@ -87,18 +88,27 @@ async def get_emotion(conn, text):
emotion = EMOJI_MAP[char]
break
try:
- await conn.websocket.send(
- json.dumps(
- {
- "type": "llm",
- "text": emoji,
- "emotion": emotion,
- "session_id": conn.session_id,
- }
- )
+ message = json.dumps(
+ {
+ "type": "llm",
+ "text": emoji,
+ "emotion": emotion,
+ "session_id": conn.session_id,
+ }
)
+
+ # 使用transport接口发送消息
+ if hasattr(conn, 'transport') and conn.transport:
+ await conn.transport.send(message)
+ elif hasattr(conn, 'websocket') and conn.websocket:
+ # 兼容旧版本
+ await conn.websocket.send(message)
+ else:
+ raise AttributeError("无法找到可用的传输层接口")
+
except Exception as e:
- conn.logger.bind(tag=TAG).warning(f"发送情绪表情失败,错误:{e}")
+ logger = setup_logging()
+ logger.warning(f"发送情绪表情失败,错误:{e}")
return
diff --git a/main/xiaozhi-server/core/websocket_server_new.py b/main/xiaozhi-server/core/websocket_server_new.py
new file mode 100644
index 00000000..238e8ba9
--- /dev/null
+++ b/main/xiaozhi-server/core/websocket_server_new.py
@@ -0,0 +1,186 @@
+import asyncio
+import websockets
+from typing import Dict, Any
+from config.logger import setup_logging
+from core.services.connection_service import ConnectionService
+from core.transport.websocket_transport import WebSocketTransport
+from config.config_loader import get_config_from_api
+from core.utils.util import check_vad_update, check_asr_update
+
+logger = setup_logging()
+
+
+class NewWebSocketServer:
+ """
+ 新的WebSocket服务器:使用新架构替代旧的ConnectionHandler
+ 集成ConnectionService、MessageRouter和新的Processor架构
+ """
+
+ def __init__(self, config: Dict[str, Any]):
+ self.config = config
+ self.logger = setup_logging()
+ self.config_lock = asyncio.Lock()
+
+ # 创建连接服务
+ self.connection_service = ConnectionService(config)
+
+ # 活跃连接管理
+ self.active_connections = set()
+
+ async def start(self):
+ """启动WebSocket服务器"""
+ server_config = self.config["server"]
+ host = server_config.get("ip", "0.0.0.0")
+ port = int(server_config.get("port", 8000))
+
+ logger.info(f"启动新架构WebSocket服务器: {host}:{port}")
+
+ async with websockets.serve(
+ self._handle_connection,
+ host,
+ port,
+ process_request=self._http_response
+ ):
+ logger.info("WebSocket服务器启动成功")
+ await asyncio.Future() # 保持服务器运行
+
+ async def _handle_connection(self, websocket):
+ """处理新连接 - 使用新架构"""
+ # 提取连接头信息
+ headers = self._extract_headers(websocket)
+
+ # 创建WebSocket传输层
+ transport = WebSocketTransport(websocket)
+
+ # 记录活跃连接
+ self.active_connections.add(transport)
+
+ try:
+ logger.info(f"新连接建立: {headers.get('device-id', 'unknown')} from {headers.get('x-real-ip', 'unknown')}")
+
+ # 使用ConnectionService处理连接
+ await self.connection_service.handle_connection(transport, headers)
+
+ except websockets.exceptions.ConnectionClosed:
+ logger.info("WebSocket连接正常关闭")
+ except Exception as e:
+ logger.error(f"处理WebSocket连接时出错: {e}", exc_info=True)
+ finally:
+ # 确保从活动连接集合中移除
+ self.active_connections.discard(transport)
+
+ # 强制关闭连接(如果还没有关闭的话)
+ try:
+ if hasattr(websocket, "closed") and not websocket.closed:
+ await websocket.close()
+ elif hasattr(websocket, "state") and websocket.state.name != "CLOSED":
+ await websocket.close()
+ except Exception as close_error:
+ logger.error(f"强制关闭WebSocket连接时出错: {close_error}")
+
+ def _extract_headers(self, websocket) -> Dict[str, str]:
+ """从WebSocket请求中提取头信息"""
+ headers = {}
+
+ # 提取请求头
+ if hasattr(websocket, 'request_headers'):
+ for name, value in websocket.request_headers.items():
+ headers[name.lower()] = value
+
+ # 提取路径参数(如果有的话)
+ if hasattr(websocket, 'path'):
+ # 可以从路径中提取device-id等参数
+ # 例如: /ws?device-id=xxx&client-id=yyy
+ from urllib.parse import urlparse, parse_qs
+ parsed = urlparse(websocket.path)
+ query_params = parse_qs(parsed.query)
+
+ for key, values in query_params.items():
+ if values:
+ headers[key] = values[0]
+
+ # 提取远程地址
+ if hasattr(websocket, 'remote_address'):
+ headers['x-real-ip'] = websocket.remote_address[0]
+
+ return headers
+
+ async def _http_response(self, websocket, request_headers):
+ """处理HTTP请求"""
+ # 检查是否为 WebSocket 升级请求
+ if request_headers.headers.get("connection", "").lower() == "upgrade":
+ # 如果是 WebSocket 请求,返回 None 允许握手继续
+ return None
+ else:
+ # 如果是普通 HTTP 请求,返回服务器状态
+ return websocket.respond(200, "New Architecture WebSocket Server is running\n")
+
+ async def update_config(self) -> bool:
+ """
+ 更新服务器配置并重新初始化组件
+
+ Returns:
+ bool: 更新是否成功
+ """
+ try:
+ async with self.config_lock:
+ logger.info("开始更新服务器配置")
+
+ # 重新获取配置
+ new_config = get_config_from_api(self.config)
+ if new_config is None:
+ logger.error("获取新配置失败")
+ return False
+
+ logger.info("获取新配置成功")
+
+ # 检查配置变化
+ config_changed = self._check_config_changes(self.config, new_config)
+
+ # 更新配置
+ self.config = new_config
+
+ # 重新创建连接服务(如果配置有重大变化)
+ if config_changed:
+ logger.info("配置有重大变化,重新创建连接服务")
+ self.connection_service = ConnectionService(new_config)
+
+ logger.info("配置更新任务执行完毕")
+ return True
+
+ except Exception as e:
+ logger.error(f"更新服务器配置失败: {str(e)}", exc_info=True)
+ return False
+
+ def _check_config_changes(self, old_config: Dict[str, Any], new_config: Dict[str, Any]) -> bool:
+ """检查配置是否有重大变化"""
+ # 检查关键配置项是否变化
+ key_configs = [
+ "selected_module",
+ "vad",
+ "asr",
+ "llm",
+ "tts",
+ "memory",
+ "intent"
+ ]
+
+ for key in key_configs:
+ if old_config.get(key) != new_config.get(key):
+ logger.info(f"配置项 {key} 发生变化")
+ return True
+
+ return False
+
+ def get_active_connections_count(self) -> int:
+ """获取活跃连接数"""
+ return len(self.active_connections)
+
+ def get_server_status(self) -> Dict[str, Any]:
+ """获取服务器状态"""
+ return {
+ "active_connections": self.get_active_connections_count(),
+ "server_type": "new_architecture",
+ "processors": self.connection_service.message_router.list_processors()
+ }
+
diff --git a/main/xiaozhi-server/core/xiaozhi_server_facade.py b/main/xiaozhi-server/core/xiaozhi_server_facade.py
new file mode 100644
index 00000000..40d84516
--- /dev/null
+++ b/main/xiaozhi-server/core/xiaozhi_server_facade.py
@@ -0,0 +1,291 @@
+#!/usr/bin/env python3
+"""
+小智服务器门面类
+统一管理所有协议服务器的启动和停止
+"""
+
+import asyncio
+from typing import Dict, Any, Optional
+from config.logger import setup_logging
+from config.config_loader import get_protocol_config, get_mqtt_server_config
+from core.servers.multi_protocol_server import MultiProtocolServer
+
+logger = setup_logging()
+
+
+class XiaozhiServerFacade:
+ """
+ 小智服务器门面类
+ 提供统一的服务器管理接口,屏蔽内部协议复杂性
+ """
+
+ def __init__(self, config: Dict[str, Any]):
+ """
+ 初始化服务器门面
+
+ Args:
+ config: 服务器配置字典
+ """
+ self.config = config
+ self.multi_protocol_server: Optional[MultiProtocolServer] = None
+ self.is_initialized = False
+ self.is_running = False
+
+ # 处理协议配置
+ self._setup_protocol_config()
+
+ def _setup_protocol_config(self):
+ """设置协议配置"""
+ try:
+ # 获取协议配置
+ try:
+ protocol_config = get_protocol_config()
+ mqtt_config = get_mqtt_server_config()
+ except Exception as e:
+ logger.warning(f"获取协议配置失败,使用默认配置: {e}")
+ # 使用默认配置
+ protocol_config = type('ProtocolConfig', (), {
+ 'websocket_enabled': True,
+ 'mqtt_enabled': False
+ })()
+ mqtt_config = type('MQTTConfig', (), {
+ 'enabled': False,
+ 'host': '0.0.0.0',
+ 'port': 1883,
+ 'udp_port': 1883,
+ 'public_ip': 'localhost',
+ 'max_connections': 1000,
+ 'heartbeat_interval': 30,
+ 'max_payload_size': 8192
+ })()
+
+ # 确定启用的协议
+ enabled_protocols = []
+
+ # WebSocket协议(默认启用)
+ if getattr(protocol_config, 'websocket_enabled', True):
+ enabled_protocols.append('websocket')
+ logger.info("WebSocket协议已启用")
+
+ # MQTT协议
+ mqtt_enabled = (
+ getattr(protocol_config, 'mqtt_enabled', False) or
+ getattr(mqtt_config, 'enabled', False)
+ )
+ if mqtt_enabled:
+ enabled_protocols.append('mqtt')
+ logger.info("MQTT协议已启用")
+
+ # 如果没有启用任何协议,默认启用WebSocket
+ if not enabled_protocols:
+ logger.warning("没有启用任何协议,默认启用WebSocket")
+ enabled_protocols = ['websocket']
+
+ # 更新配置
+ self.config['enabled_protocols'] = enabled_protocols
+
+ # 添加MQTT服务器配置
+ self.config['mqtt_server'] = {
+ 'enabled': getattr(mqtt_config, 'enabled', False),
+ 'host': getattr(mqtt_config, 'host', '0.0.0.0'),
+ 'port': getattr(mqtt_config, 'port', 1883),
+ 'udp_port': getattr(mqtt_config, 'udp_port', 1883),
+ 'public_ip': getattr(mqtt_config, 'public_ip', 'localhost'),
+ 'max_connections': getattr(mqtt_config, 'max_connections', 1000),
+ 'heartbeat_interval': getattr(mqtt_config, 'heartbeat_interval', 30),
+ 'max_payload_size': getattr(mqtt_config, 'max_payload_size', 8192)
+ }
+
+ logger.info(f"启用的协议: {enabled_protocols}")
+
+ except Exception as e:
+ logger.error(f"设置协议配置失败: {e}")
+ # 使用最基本的配置
+ self.config['enabled_protocols'] = ['websocket']
+ self.config['mqtt_server'] = {
+ 'enabled': False,
+ 'host': '0.0.0.0',
+ 'port': 1883,
+ 'udp_port': 1883,
+ 'public_ip': 'localhost',
+ 'max_connections': 1000,
+ 'heartbeat_interval': 30,
+ 'max_payload_size': 8192
+ }
+
+ async def initialize(self):
+ """初始化服务器"""
+ if self.is_initialized:
+ logger.warning("服务器已经初始化")
+ return
+
+ try:
+ logger.info("正在初始化小智服务器...")
+
+ # 创建多协议服务器
+ self.multi_protocol_server = MultiProtocolServer(self.config)
+
+ self.is_initialized = True
+ logger.info("小智服务器初始化完成")
+
+ except Exception as e:
+ logger.error(f"初始化服务器失败: {e}")
+ raise
+
+ async def start(self):
+ """启动服务器"""
+ if not self.is_initialized:
+ await self.initialize()
+
+ if self.is_running:
+ logger.warning("服务器已经在运行中")
+ return
+
+ try:
+ logger.info("正在启动小智服务器...")
+
+ # 启动多协议服务器
+ await self.multi_protocol_server.start()
+
+ self.is_running = True
+ logger.info("小智服务器启动成功")
+
+ except Exception as e:
+ logger.error(f"启动服务器失败: {e}")
+ self.is_running = False
+ raise
+
+ async def stop(self):
+ """停止服务器"""
+ if not self.is_running:
+ logger.info("服务器未在运行")
+ return
+
+ try:
+ logger.info("正在停止小智服务器...")
+
+ if self.multi_protocol_server:
+ await self.multi_protocol_server.stop()
+
+ self.is_running = False
+ logger.info("小智服务器已停止")
+
+ except Exception as e:
+ logger.error(f"停止服务器失败: {e}")
+
+ async def restart(self):
+ """重启服务器"""
+ logger.info("重启小智服务器...")
+ await self.stop()
+ await asyncio.sleep(1) # 等待清理完成
+ await self.start()
+
+ async def update_config(self, new_config: Dict[str, Any]) -> bool:
+ """
+ 更新服务器配置
+
+ Args:
+ new_config: 新的配置字典
+
+ Returns:
+ bool: 更新是否成功
+ """
+ try:
+ logger.info("更新服务器配置...")
+
+ # 更新配置
+ self.config.update(new_config)
+ self._setup_protocol_config()
+
+ # 如果服务器正在运行,更新多协议服务器配置
+ if self.is_running and self.multi_protocol_server:
+ success = await self.multi_protocol_server.update_config(self.config)
+ if success:
+ logger.info("服务器配置更新成功")
+ else:
+ logger.error("服务器配置更新失败")
+ return success
+
+ logger.info("配置更新完成(服务器未运行)")
+ return True
+
+ except Exception as e:
+ logger.error(f"更新配置失败: {e}")
+ return False
+
+ def get_server_status(self) -> Dict[str, Any]:
+ """获取服务器状态"""
+ base_status = {
+ 'is_initialized': self.is_initialized,
+ 'is_running': self.is_running,
+ 'enabled_protocols': self.config.get('enabled_protocols', [])
+ }
+
+ if self.multi_protocol_server:
+ server_status = self.multi_protocol_server.get_server_status()
+ base_status.update(server_status)
+
+ return base_status
+
+ def get_active_connections_count(self) -> Dict[str, int]:
+ """获取各协议的活跃连接数"""
+ if self.multi_protocol_server:
+ return self.multi_protocol_server.get_active_connections_count()
+ return {}
+
+ def get_supported_protocols(self) -> list:
+ """获取支持的协议列表"""
+ if self.multi_protocol_server:
+ return self.multi_protocol_server.get_supported_protocols()
+ return ['websocket', 'mqtt']
+
+ def is_protocol_enabled(self, protocol: str) -> bool:
+ """检查协议是否启用"""
+ enabled_protocols = self.config.get('enabled_protocols', [])
+ return protocol in enabled_protocols
+
+ async def broadcast_message(self, message: Dict[str, Any], protocol: Optional[str] = None):
+ """
+ 向所有连接广播消息
+
+ Args:
+ message: 要广播的消息
+ protocol: 指定协议,None表示向所有协议广播
+ """
+ if self.multi_protocol_server:
+ await self.multi_protocol_server.broadcast_message(message, protocol)
+
+ def get_websocket_info(self) -> Dict[str, Any]:
+ """获取WebSocket连接信息"""
+ if not self.is_protocol_enabled('websocket'):
+ return {'enabled': False}
+
+ server_config = self.config.get('server', {})
+ return {
+ 'enabled': True,
+ 'host': server_config.get('ip', '0.0.0.0'),
+ 'port': server_config.get('port', 8000),
+ 'path': '/xiaozhi/v1/'
+ }
+
+ def get_mqtt_info(self) -> Dict[str, Any]:
+ """获取MQTT连接信息"""
+ if not self.is_protocol_enabled('mqtt'):
+ return {'enabled': False}
+
+ mqtt_config = self.config.get('mqtt_server', {})
+ return {
+ 'enabled': True,
+ 'host': mqtt_config.get('host', '0.0.0.0'),
+ 'port': mqtt_config.get('port', 1883),
+ 'udp_port': mqtt_config.get('udp_port', 1883),
+ 'public_ip': mqtt_config.get('public_ip', 'localhost')
+ }
+
+ def get_connection_info(self) -> Dict[str, Any]:
+ """获取所有协议的连接信息"""
+ return {
+ 'websocket': self.get_websocket_info(),
+ 'mqtt': self.get_mqtt_info(),
+ 'active_connections': self.get_active_connections_count()
+ }
diff --git a/main/xiaozhi-server/plugins_func/functions/play_music.py b/main/xiaozhi-server/plugins_func/functions/play_music.py
index 2cbc4018..4bad3db0 100644
--- a/main/xiaozhi-server/plugins_func/functions/play_music.py
+++ b/main/xiaozhi-server/plugins_func/functions/play_music.py
@@ -5,7 +5,7 @@ import random
import difflib
import traceback
from pathlib import Path
-from core.handle.sendAudioHandle import send_stt_message
+# from core.handle.sendAudioHandle import send_stt_message # 未使用,已移除
from plugins_func.register import register_function, ToolType, ActionResponse, Action
from core.utils.dialogue import Message
from core.providers.tts.dto.dto import TTSMessageDTO, SentenceType, ContentType