""" 工具注册表 - 支持动态工具发现和加载 提供插件化架构,允许通过配置动态加载和替换工具。 """ import logging from typing import Any, Callable, Dict, List, Optional, Type logger = logging.getLogger(__name__) class ToolRegistry: """工具注册表 - 管理所有可用工具的注册和获取""" def __init__(self): self._tools: Dict[str, Any] = {} self._tool_info: Dict[str, Dict] = {} def register(self, name: str, tool_class: Any, metadata: Optional[Dict] = None) -> None: """ 注册工具到注册表 Args: name: 工具名称(唯一标识) tool_class: 工具类或工厂函数 metadata: 工具元数据(描述、版本、作者等) """ if name in self._tools: logger.warning(f"Tool '{name}' already registered, overwriting") self._tools[name] = tool_class self._tool_info[name] = metadata or {} logger.info(f"Registered tool: {name}") def get(self, name: str) -> Any: """ 获取已注册的工具 Args: name: 工具名称 Returns: 工具类或工厂函数 Raises: KeyError: 工具未注册 """ if name not in self._tools: raise KeyError(f"Tool '{name}' not registered. Available tools: {list(self._tools.keys())}") return self._tools[name] def get_instance(self, name: str, *args, **kwargs) -> Any: """ 获取工具实例(调用注册的类或工厂函数) Args: name: 工具名称 *args: 位置参数 **kwargs: 关键字参数 Returns: 工具实例 """ tool = self.get(name) if callable(tool): return tool(*args, **kwargs) return tool def list_tools(self) -> List[str]: """列出所有已注册的工具名称""" return list(self._tools.keys()) def get_info(self, name: str) -> Dict: """获取工具元数据""" return self._tool_info.get(name, {}) def unregister(self, name: str) -> bool: """ 注销工具 Args: name: 工具名称 Returns: 是否成功注销 """ if name in self._tools: del self._tools[name] del self._tool_info[name] logger.info(f"Unregistered tool: {name}") return True return False def has(self, name: str) -> bool: """检查工具是否已注册""" return name in self._tools def clear(self) -> None: """清除所有已注册的工具""" self._tools.clear() self._tool_info.clear() logger.info("Cleared all registered tools") # 全局注册表实例 _global_registry: Optional[ToolRegistry] = None def get_registry() -> ToolRegistry: """获取全局工具注册表实例""" global _global_registry if _global_registry is None: _global_registry = ToolRegistry() _register_default_tools() return _global_registry def _register_default_tools() -> None: """注册默认工具""" global _global_registry # 延迟导入避免循环依赖 try: from runners import CobolRunner, NativeJavaRunner, SparkJavaRunner from runners.gixsql_runner import GixsqlCobolRunner from agents.llm import LLMClient from comparator import FieldComparator _global_registry.register("cobol_runner", CobolRunner, {"type": "runner", "description": "COBOL compiler and runner"}) _global_registry.register("java_runner", NativeJavaRunner, {"type": "runner", "description": "Java native runner"}) _global_registry.register("spark_runner", SparkJavaRunner, {"type": "runner", "description": "Spark Java runner"}) _global_registry.register("gixsql_runner", GixsqlCobolRunner, {"type": "runner", "description": "DB COBOL runner with gixsql"}) _global_registry.register("llm_client", LLMClient, {"type": "agent", "description": "LLM API client"}) _global_registry.register("comparator", FieldComparator, {"type": "comparator", "description": "Field-level comparison"}) logger.info(f"Registered {len(_global_registry.list_tools())} default tools") except ImportError as e: logger.warning(f"Failed to register default tools: {e}") class ToolConfig: """工具配置 - 从配置文件或环境变量加载工具配置""" @staticmethod def get_runner_mode(config: Dict) -> str: """获取运行器模式""" return config.get("runner_mode", "native") @staticmethod def get_runner_class(registry: ToolRegistry, mode: str) -> Any: """根据模式获取运行器类""" runner_map = { "native": "java_runner", "spark": "spark_runner", "cobol": "cobol_runner", } tool_name = runner_map.get(mode, "java_runner") return registry.get(tool_name) @staticmethod def get_llm_config(config: Dict) -> Dict: """获取LLM配置""" return { "model": config.get("llm_model", None), "timeout": config.get("llm_timeout", 15), "cache_dir": config.get("llm_cache_dir", ".cache/llm"), }