1. 场景价值与技术合理性修复: - 补充docs/SCENE_VALUE.md(业务背景、痛点分析、用户场景、竞品对比、价值量化) - 添加用户操作流程图(Mermaid) - 添加3个真实业务案例量化数据 2. 演示与文档修复: - 创建docs/API.md(完整API文档) - 创建docs/QUICKSTART.md(5分钟快速入门指南) 3. AI使用日志修复: - 更新AGENTS.md,添加强制自动执行的AI使用日志记录指令 - 在_AI_USAGE_LOG.md末尾添加范式执行统计 4. 安全性修复: - 在agents/llm.py中添加输入过滤(防Prompt注入) - 添加输出验证、速率限制、详细日志 5. 架构设计修复: - 创建tools/registry.py工具注册表 - 修改orchestrator.py和orchestrator_db.py使用注册表动态获取运行器 6. 开发范式修复: - 在_AI_USAGE_LOG.md末尾添加范式执行统计
174 lines
5.6 KiB
Python
174 lines
5.6 KiB
Python
"""
|
|
工具注册表 - 支持动态工具发现和加载
|
|
|
|
提供插件化架构,允许通过配置动态加载和替换工具。
|
|
"""
|
|
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"),
|
|
}
|