mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-30 05:13:59 +08:00
Merge pull request #3304 from MsThink/main
fix(plugins): 展开模块级别插件名为具体函数名
This commit is contained in:
@@ -35,7 +35,7 @@ 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 core.providers.tools.unified_tool_handler import UnifiedToolHandler
|
||||||
from plugins_func.loadplugins import auto_import_modules
|
from plugins_func.loadplugins import auto_import_modules
|
||||||
from plugins_func.register import Action, ActionResponse
|
from plugins_func.register import Action, ActionResponse, all_function_registry, module_func_map
|
||||||
from core.auth import AuthenticationError
|
from core.auth import AuthenticationError
|
||||||
from config.config_loader import get_private_config_from_api
|
from config.config_loader import get_private_config_from_api
|
||||||
from core.providers.tts.dto.dto import ContentType, TTSMessageDTO, SentenceType
|
from core.providers.tts.dto.dto import ContentType, TTSMessageDTO, SentenceType
|
||||||
@@ -878,10 +878,27 @@ class ConnectionHandler:
|
|||||||
plugin_from_server = private_config.get("plugins", {})
|
plugin_from_server = private_config.get("plugins", {})
|
||||||
for plugin, config_str in plugin_from_server.items():
|
for plugin, config_str in plugin_from_server.items():
|
||||||
plugin_from_server[plugin] = json.loads(config_str)
|
plugin_from_server[plugin] = json.loads(config_str)
|
||||||
|
# 将模块级别的插件配置复制到各具体函数名,
|
||||||
|
# 方便后续 per-function 查找描述、news_sources 等配置
|
||||||
|
for module_name, func_names in module_func_map.items():
|
||||||
|
if module_name in plugin_from_server:
|
||||||
|
module_config = plugin_from_server[module_name]
|
||||||
|
for func_name in func_names:
|
||||||
|
if func_name not in plugin_from_server:
|
||||||
|
plugin_from_server[func_name] = module_config
|
||||||
self.config["plugins"] = plugin_from_server
|
self.config["plugins"] = plugin_from_server
|
||||||
|
# 将模块级别的插件名展开为具体函数名
|
||||||
|
expanded_functions = []
|
||||||
|
for plugin_key in plugin_from_server.keys():
|
||||||
|
if plugin_key in all_function_registry:
|
||||||
|
expanded_functions.append(plugin_key)
|
||||||
|
elif plugin_key in module_func_map:
|
||||||
|
expanded_functions.extend(module_func_map[plugin_key])
|
||||||
|
else:
|
||||||
|
expanded_functions.append(plugin_key)
|
||||||
self.config["Intent"][self.config["selected_module"]["Intent"]][
|
self.config["Intent"][self.config["selected_module"]["Intent"]][
|
||||||
"functions"
|
"functions"
|
||||||
] = plugin_from_server.keys()
|
] = expanded_functions
|
||||||
if private_config.get("prompt", None) is not None:
|
if private_config.get("prompt", None) is not None:
|
||||||
self.config["prompt"] = private_config["prompt"]
|
self.config["prompt"] = private_config["prompt"]
|
||||||
# 获取声纹信息
|
# 获取声纹信息
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ from typing import Dict, Any, TYPE_CHECKING
|
|||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from core.connection import ConnectionHandler
|
from core.connection import ConnectionHandler
|
||||||
from ..base import ToolType, ToolDefinition, ToolExecutor
|
from ..base import ToolType, ToolDefinition, ToolExecutor
|
||||||
from plugins_func.register import all_function_registry, Action, ActionResponse
|
from plugins_func.register import all_function_registry, module_func_map, Action, ActionResponse
|
||||||
|
|
||||||
|
|
||||||
class ServerPluginExecutor(ToolExecutor):
|
class ServerPluginExecutor(ToolExecutor):
|
||||||
@@ -54,6 +54,43 @@ class ServerPluginExecutor(ToolExecutor):
|
|||||||
response=str(e),
|
response=str(e),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _expand_plugin_names(self, config_functions):
|
||||||
|
"""将模块级别的插件名展开为具体函数名。
|
||||||
|
|
||||||
|
在一个 function 文件中可能注册多个 @register_function,
|
||||||
|
配置中如果使用的是模块名(文件名),需要展开为具体的函数名列表。
|
||||||
|
"""
|
||||||
|
if not isinstance(config_functions, list):
|
||||||
|
try:
|
||||||
|
config_functions = list(config_functions)
|
||||||
|
except TypeError:
|
||||||
|
return []
|
||||||
|
|
||||||
|
expanded = []
|
||||||
|
for name in config_functions:
|
||||||
|
if name in all_function_registry:
|
||||||
|
# 精确匹配函数名,直接保留
|
||||||
|
expanded.append(name)
|
||||||
|
elif name in module_func_map:
|
||||||
|
# 模块名,展开为该模块下所有注册函数名
|
||||||
|
expanded.extend(module_func_map[name])
|
||||||
|
else:
|
||||||
|
# 未知名称,保留原值(可能是 MCP 或其他工具)
|
||||||
|
expanded.append(name)
|
||||||
|
return expanded
|
||||||
|
|
||||||
|
def _get_plugin_description(self, func_name):
|
||||||
|
"""获取插件函数的描述,优先精确匹配函数名,其次匹配模块名。"""
|
||||||
|
plugins = self.config.get("plugins", {})
|
||||||
|
# 精确匹配函数名
|
||||||
|
if func_name in plugins:
|
||||||
|
return plugins[func_name].get("description", "")
|
||||||
|
# 通过 module_func_map 反向查找模块名
|
||||||
|
for module_name, func_names in module_func_map.items():
|
||||||
|
if func_name in func_names and module_name in plugins:
|
||||||
|
return plugins[module_name].get("description", "")
|
||||||
|
return ""
|
||||||
|
|
||||||
def get_tools(self) -> Dict[str, ToolDefinition]:
|
def get_tools(self) -> Dict[str, ToolDefinition]:
|
||||||
"""获取所有注册的服务端插件工具"""
|
"""获取所有注册的服务端插件工具"""
|
||||||
tools = {}
|
tools = {}
|
||||||
@@ -66,12 +103,8 @@ class ServerPluginExecutor(ToolExecutor):
|
|||||||
self.config["selected_module"]["Intent"]
|
self.config["selected_module"]["Intent"]
|
||||||
].get("functions", [])
|
].get("functions", [])
|
||||||
|
|
||||||
# 转换为列表
|
# 将模块级别的插件名展开为具体函数名(兜底机制)
|
||||||
if not isinstance(config_functions, list):
|
config_functions = self._expand_plugin_names(config_functions)
|
||||||
try:
|
|
||||||
config_functions = list(config_functions)
|
|
||||||
except TypeError:
|
|
||||||
config_functions = []
|
|
||||||
|
|
||||||
# 合并所有需要的函数
|
# 合并所有需要的函数
|
||||||
all_required_functions = list(set(necessary_functions + config_functions))
|
all_required_functions = list(set(necessary_functions + config_functions))
|
||||||
@@ -79,12 +112,8 @@ class ServerPluginExecutor(ToolExecutor):
|
|||||||
for func_name in all_required_functions:
|
for func_name in all_required_functions:
|
||||||
func_item = all_function_registry.get(func_name)
|
func_item = all_function_registry.get(func_name)
|
||||||
if func_item:
|
if func_item:
|
||||||
# 从函数注册中获取描述
|
# 从函数注册中获取描述(支持模块名和函数名的双层查找)
|
||||||
fun_description = (
|
fun_description = self._get_plugin_description(func_name)
|
||||||
self.config.get("plugins", {})
|
|
||||||
.get(func_name, {})
|
|
||||||
.get("description", "")
|
|
||||||
)
|
|
||||||
if fun_description is not None and len(fun_description) > 0:
|
if fun_description is not None and len(fun_description) > 0:
|
||||||
if "function" in func_item.description and isinstance(
|
if "function" in func_item.description and isinstance(
|
||||||
func_item.description["function"], dict
|
func_item.description["function"], dict
|
||||||
|
|||||||
@@ -78,6 +78,8 @@ class DeviceTypeRegistry:
|
|||||||
|
|
||||||
# 初始化函数注册字典
|
# 初始化函数注册字典
|
||||||
all_function_registry = {}
|
all_function_registry = {}
|
||||||
|
# 模块名 -> 函数名列表的映射,用于将模块级别的插件名展开为具体的函数名
|
||||||
|
module_func_map = {}
|
||||||
|
|
||||||
|
|
||||||
def register_function(name, desc, type=None):
|
def register_function(name, desc, type=None):
|
||||||
@@ -85,6 +87,9 @@ def register_function(name, desc, type=None):
|
|||||||
|
|
||||||
def decorator(func):
|
def decorator(func):
|
||||||
all_function_registry[name] = FunctionItem(name, desc, func, type)
|
all_function_registry[name] = FunctionItem(name, desc, func, type)
|
||||||
|
# 记录模块名到函数名的映射,用于 expand 模块级别的插件配置
|
||||||
|
module_name = func.__module__.split(".")[-1]
|
||||||
|
module_func_map.setdefault(module_name, []).append(name)
|
||||||
logger.bind(tag=TAG).debug(f"函数 '{name}' 已加载,可以注册使用")
|
logger.bind(tag=TAG).debug(f"函数 '{name}' 已加载,可以注册使用")
|
||||||
return func
|
return func
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user