Files
hangshuo652 b94757d9df feat: V3系统评审问题修复
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末尾添加范式执行统计
2026-08-29 13:23:28 +08:00

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"),
}