Files
agent_management/template_manager.py
T

568 lines
20 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
Template Manager - 动态管理 Agent 模板
支持通过 API 添加、更新、删除模板,无需修改代码或重新部署
"""
import logging
from typing import Dict, List, Optional, Any
from datetime import datetime
from sqlalchemy.orm import Session
from sqlalchemy import or_
from database import Template, AgentType, SessionLocal
logger = logging.getLogger(__name__)
# 默认模板配置(用于初始化数据库)
DEFAULT_TEMPLATES = {
"echo_agent": {
"display_name": "Echo Agent",
"description": "简单的回显测试 Agent",
"image": "agnettaiji.azurecr.io/ai-agents/echo-agent:latest",
"port": 8000,
"agent_framework": "api",
"env_requirements": {},
},
"search_agent": {
"display_name": "Search Agent",
"description": "搜索 Agent (LangChain)",
"image": "agnettaiji.azurecr.io/ai-agents/search-agent:latest",
"port": 8080,
"agent_framework": "langchain",
"env_requirements": {},
},
"search_agent_a2a": {
"display_name": "Search Agent A2A",
"description": "搜索 Agent (A2A 协议)",
"image": "agnettaiji.azurecr.io/ai-agents/search-agent-a2a:latest",
"port": 8080,
"agent_framework": "a2a",
"env_requirements": {},
},
"search_agent_mcp": {
"display_name": "Search Agent MCP",
"description": "搜索 Agent (MCP 协议)",
"image": "agnettaiji.azurecr.io/ai-agents/search-agent-mcp:latest",
"port": 8080,
"agent_framework": "mcp",
"env_requirements": {},
},
"mysql_agent": {
"display_name": "MySQL Agent",
"description": "MySQL 数据库 Agent",
"image": "agnettaiji.azurecr.io/ai-agents/mysql-agent:latest",
"port": 8000,
"agent_framework": "api",
"env_requirements": {
"required": {
"MYSQL_HOST": "MySQL 主机地址",
"MYSQL_USER": "MySQL 用户名",
"MYSQL_PASSWORD": "MySQL 密码",
"MYSQL_DATABASE": "MySQL 数据库名",
}
},
},
"postgresql_agent": {
"display_name": "PostgreSQL Agent",
"description": "PostgreSQL 数据库 Agent",
"image": "agnettaiji.azurecr.io/ai-agents/postgresql-agent:latest",
"port": 8000,
"agent_framework": "api",
"env_requirements": {
"required": {
"POSTGRES_HOST": "PostgreSQL 主机地址",
"POSTGRES_USER": "PostgreSQL 用户名",
"POSTGRES_PASSWORD": "PostgreSQL 密码",
"POSTGRES_DATABASE": "PostgreSQL 数据库名",
}
},
},
"jina_search_agent": {
"display_name": "Jina Search Agent",
"description": "Jina AI 搜索 Agent",
"image": "agnettaiji.azurecr.io/ai-agents/jina-search-agent:latest",
"port": 8000,
"agent_framework": "api",
"env_requirements": {
"required": {
"JINA_API_KEY": "Jina API 密钥",
}
},
},
"azure_blob_agent": {
"display_name": "Azure Blob Agent",
"description": "Azure Blob 存储 Agent",
"image": "agnettaiji.azurecr.io/ai-agents/azure-blob-agent:latest",
"port": 8000,
"agent_framework": "api",
"env_requirements": {
"required": {
"AZURE_STORAGE_CONNECTION_STRING": "Azure 存储连接字符串",
}
},
},
"azure_blob_agent_mcp": {
"display_name": "Azure Blob Agent MCP",
"description": "Azure Blob 存储 Agent (MCP 协议)",
"image": "agnettaiji.azurecr.io/ai-agents/azure-blob-agent-mcp:latest",
"port": 8000,
"agent_framework": "mcp",
"env_requirements": {
"required": {
"AZURE_STORAGE_CONNECTION_STRING": "Azure 存储连接字符串",
}
},
},
"azure_blob_agent_a2a": {
"display_name": "Azure Blob Agent A2A",
"description": "Azure Blob 存储 Agent (A2A 协议)",
"image": "agnettaiji.azurecr.io/ai-agents/azure-blob-agent-a2a:latest",
"port": 8000,
"agent_framework": "a2a",
"env_requirements": {
"required": {
"AZURE_STORAGE_CONNECTION_STRING": "Azure 存储连接字符串",
}
},
},
"a2a_litellm_agent": {
"display_name": "A2A LiteLLM Agent",
"description": "LiteLLM A2A 协议 Agent",
"image": "agnettaiji.azurecr.io/ai-agents/a2a-litellm-agent:latest",
"port": 8000,
"agent_framework": "a2a",
"env_requirements": {},
},
"coding_a2a_agent": {
"display_name": "Coding A2A Agent",
"description": "Claude Code 风格编程 Agent(Pydantic AI + A2A)",
"image": "agnettaiji.azurecr.io/ai-agents/coding-a2a-agent:latest",
"port": 8000,
"agent_framework": "a2a",
"env_requirements": {
"required": {
"OPENAI_BASE_URL": "LiteLLM / OpenAI 兼容网关地址",
},
"optional": {
"OPENAI_API_KEY": "模型 API Key",
"LITELLM_API_KEY": "模型 API Key(兼容变量)",
"MODEL_NAME": "模型名称",
"LITELLM_MODEL": "模型名称(兼容变量)",
"WORK_DIR": "工作区目录,默认 /workspace",
"AGENT_ROLE_NAME": "启动时指定角色名称,例如 backend / reviewer / planner",
"AGENT_INSTRUCTION_TEXT": "启动时注入的角色/行为说明文本,支持类似 AGENTS.md / claude.md 内容",
"AGENT_INSTRUCTION_FILE": "启动时读取的角色说明文件路径,内容会并入系统提示词",
"SERVICE_PORT": "服务端口,默认 8000",
},
},
},
"code_ai_agent": {
"display_name": "Code AI Agent",
"description": "代码 AI 助手 Agent",
"image": "agnettaiji.azurecr.io/ai-agents/code-ai-agent:latest",
"port": 8000,
"agent_framework": "api",
"env_requirements": {},
},
"facebook_agent": {
"display_name": "Facebook Agent",
"description": "Facebook 社交媒体 Agent",
"image": "agnettaiji.azurecr.io/ai-agents/facebook-agent:latest",
"port": 8000,
"agent_framework": "api",
"env_requirements": {},
},
"media_downloader": {
"display_name": "Media Downloader",
"description": "媒体下载 Agent(YouTube、视频等)",
"image": "agnettaiji.azurecr.io/ai-agents/media-downloader:latest",
"port": 8000,
"agent_framework": "api",
"env_requirements": {},
},
"content_analyzer": {
"display_name": "Content Analyzer",
"description": "内容分析 Agent",
"image": "agnettaiji.azurecr.io/ai-agents/content-analyzer:latest",
"port": 8000,
"agent_framework": "api",
"env_requirements": {},
},
"huoke": {
"display_name": "Huoke Agent",
"description": "获客 Agent",
"image": "agnettaiji.azurecr.io/ai-agents/huoke:latest",
"port": 8000,
"agent_framework": "api",
"env_requirements": {},
},
"microsoft_learn_agent": {
"display_name": "Microsoft Learn Agent",
"description": "Microsoft Learn 文档 Agent",
"image": "agnettaiji.azurecr.io/ai-agents/microsoft-learn-agent:latest",
"port": 8000,
"agent_framework": "mcp",
"env_requirements": {},
},
"aws_docs_mcp": {
"display_name": "AWS Docs MCP Agent",
"description": "AWS 文档 MCP Agent",
"image": "agnettaiji.azurecr.io/ai-agents/aws-docs-mcp:latest",
"port": 8000,
"agent_framework": "mcp",
"env_requirements": {},
},
"google_mcp": {
"display_name": "Google MCP Agent",
"description": "Google MCP Agent",
"image": "agnettaiji.azurecr.io/ai-agents/google-mcp:latest",
"port": 8000,
"agent_framework": "mcp",
"env_requirements": {},
},
}
class TemplateManager:
"""模板管理器 - 提供模板的 CRUD 操作和动态加载"""
_instance = None
_cache: Dict[str, Dict] = {}
_cache_time: Optional[datetime] = None
_cache_ttl = 60 # 缓存有效期(秒)
def __new__(cls):
if cls._instance is None:
cls._instance = super().__new__(cls)
cls._instance._initialized = False
return cls._instance
def __init__(self):
if self._initialized:
return
self._initialized = True
self._ensure_default_templates()
def _get_db(self) -> Session:
"""获取数据库会话"""
return SessionLocal()
def _ensure_default_templates(self):
"""确保默认模板存在于数据库中"""
db = self._get_db()
try:
for name, config in DEFAULT_TEMPLATES.items():
existing = db.query(Template).filter(Template.name == name).first()
if not existing:
template = Template(
name=name,
display_name=config.get("display_name", name),
description=config.get("description", ""),
agent_type=AgentType.PLATFORM,
agent_framework=config.get("agent_framework", "api"),
image=config["image"],
port=config.get("port", 8000),
env_requirements=config.get("env_requirements", {}),
tools_config=config.get("tools_config", {}),
is_active=True,
created_by="system",
)
db.add(template)
logger.info(f"✅ 初始化模板: {name}")
db.commit()
logger.info("✅ 默认模板初始化完成")
except Exception as e:
db.rollback()
logger.error(f"初始化默认模板失败: {e}")
finally:
db.close()
def _refresh_cache(self, db: Session = None):
"""刷新模板缓存"""
should_close = False
if db is None:
db = self._get_db()
should_close = True
try:
templates = db.query(Template).filter(Template.is_active == True).all()
self._cache = {}
for t in templates:
self._cache[t.name] = {
"id": t.id,
"name": t.name,
"template": t.name, # 兼容 MCP-Server 期望的 template 字段
"display_name": t.display_name,
"description": t.description,
"agent_type": t.agent_type.value if t.agent_type else "platform",
"agent_framework": t.agent_framework or "api",
"image": t.image,
"port": t.port or 8000,
"env_requirements": t.env_requirements or {},
"tools_config": t.tools_config or {},
"cpu_request": t.cpu_request,
"cpu_limit": t.cpu_limit,
"memory_request": t.memory_request,
"memory_limit": t.memory_limit,
"is_active": t.is_active,
"created_at": t.created_at.isoformat() if t.created_at else None,
}
self._cache_time = datetime.utcnow()
logger.debug(f"模板缓存已刷新: {len(self._cache)} 个模板")
finally:
if should_close:
db.close()
def _is_cache_valid(self) -> bool:
"""检查缓存是否有效"""
if not self._cache_time:
return False
elapsed = (datetime.utcnow() - self._cache_time).total_seconds()
return elapsed < self._cache_ttl
def get_all_templates(self, include_inactive: bool = False) -> List[Dict]:
"""获取所有模板列表"""
if not self._is_cache_valid():
self._refresh_cache()
templates = list(self._cache.values())
if not include_inactive:
templates = [t for t in templates if t.get("is_active", True)]
return templates
def get_template(self, name: str) -> Optional[Dict]:
"""根据名称获取模板"""
if not self._is_cache_valid():
self._refresh_cache()
return self._cache.get(name)
def get_template_names(self) -> List[str]:
"""获取所有活跃模板名称列表"""
if not self._is_cache_valid():
self._refresh_cache()
return [name for name, t in self._cache.items() if t.get("is_active", True)]
def get_image(self, template_name: str) -> str:
"""获取模板镜像地址"""
template = self.get_template(template_name)
if template:
return template["image"]
# 回退到默认搜索 Agent
return DEFAULT_TEMPLATES.get("search_agent", {}).get(
"image", "agnettaiji.azurecr.io/ai-agents/search-agent:latest"
)
def get_port(self, template_name: str) -> int:
"""获取模板端口"""
template = self.get_template(template_name)
if template:
return template.get("port", 8000)
return 8000
def get_env_requirements(self, template_name: str) -> Dict:
"""获取模板环境变量要求"""
template = self.get_template(template_name)
if template:
return template.get("env_requirements", {})
return {}
def template_exists(self, name: str) -> bool:
"""检查模板是否存在"""
return self.get_template(name) is not None
def create_template(
self,
name: str,
image: str,
display_name: Optional[str] = None,
description: Optional[str] = None,
port: int = 8000,
agent_framework: str = "api",
env_requirements: Optional[Dict] = None,
tools_config: Optional[Dict] = None,
cpu_request: Optional[str] = None,
cpu_limit: Optional[str] = None,
memory_request: Optional[str] = None,
memory_limit: Optional[str] = None,
created_by: str = "api",
) -> Dict:
"""创建新模板"""
db = self._get_db()
try:
# 检查是否已存在
existing = db.query(Template).filter(Template.name == name).first()
if existing:
raise ValueError(f"模板 '{name}' 已存在")
# 创建模板
template = Template(
name=name,
display_name=display_name or name,
description=description or "",
agent_type=AgentType.PLATFORM,
agent_framework=agent_framework,
image=image,
port=port,
env_requirements=env_requirements or {},
tools_config=tools_config or {},
cpu_request=cpu_request,
cpu_limit=cpu_limit,
memory_request=memory_request,
memory_limit=memory_limit,
is_active=True,
created_by=created_by,
)
db.add(template)
db.commit()
db.refresh(template)
# 刷新缓存
self._refresh_cache(db)
logger.info(f"✅ 创建模板成功: {name}")
return self.get_template(name)
except Exception as e:
db.rollback()
logger.error(f"创建模板失败: {e}")
raise
finally:
db.close()
def update_template(
self,
name: str,
image: Optional[str] = None,
display_name: Optional[str] = None,
description: Optional[str] = None,
port: Optional[int] = None,
agent_framework: Optional[str] = None,
env_requirements: Optional[Dict] = None,
tools_config: Optional[Dict] = None,
cpu_request: Optional[str] = None,
cpu_limit: Optional[str] = None,
memory_request: Optional[str] = None,
memory_limit: Optional[str] = None,
is_active: Optional[bool] = None,
) -> Dict:
"""更新模板"""
db = self._get_db()
try:
template = db.query(Template).filter(Template.name == name).first()
if not template:
raise ValueError(f"模板 '{name}' 不存在")
# 更新字段
if image is not None:
template.image = image
if display_name is not None:
template.display_name = display_name
if description is not None:
template.description = description
if port is not None:
template.port = port
if agent_framework is not None:
template.agent_framework = agent_framework
if env_requirements is not None:
template.env_requirements = env_requirements
if tools_config is not None:
template.tools_config = tools_config
if cpu_request is not None:
template.cpu_request = cpu_request
if cpu_limit is not None:
template.cpu_limit = cpu_limit
if memory_request is not None:
template.memory_request = memory_request
if memory_limit is not None:
template.memory_limit = memory_limit
if is_active is not None:
template.is_active = is_active
template.updated_at = datetime.utcnow()
db.commit()
# 刷新缓存
self._refresh_cache(db)
logger.info(f"✅ 更新模板成功: {name}")
return self.get_template(name)
except Exception as e:
db.rollback()
logger.error(f"更新模板失败: {e}")
raise
finally:
db.close()
def delete_template(self, name: str, force: bool = False) -> bool:
"""删除模板(软删除,除非 force=True)"""
db = self._get_db()
try:
template = db.query(Template).filter(Template.name == name).first()
if not template:
raise ValueError(f"模板 '{name}' 不存在")
if force:
# 硬删除
db.delete(template)
logger.info(f"✅ 永久删除模板: {name}")
else:
# 软删除
template.is_active = False
template.updated_at = datetime.utcnow()
logger.info(f"✅ 禁用模板: {name}")
db.commit()
# 刷新缓存
self._refresh_cache(db)
return True
except Exception as e:
db.rollback()
logger.error(f"删除模板失败: {e}")
raise
finally:
db.close()
def search_templates(self, keyword: str) -> List[Dict]:
"""搜索模板"""
db = self._get_db()
try:
templates = db.query(Template).filter(
Template.is_active == True,
or_(
Template.name.ilike(f"%{keyword}%"),
Template.display_name.ilike(f"%{keyword}%"),
Template.description.ilike(f"%{keyword}%"),
)
).all()
return [
{
"name": t.name,
"template": t.name, # 兼容 MCP-Server 期望的 template 字段
"display_name": t.display_name,
"description": t.description,
"image": t.image,
"port": t.port,
"agent_framework": t.agent_framework,
}
for t in templates
]
finally:
db.close()
def invalidate_cache(self):
"""手动使缓存失效"""
self._cache_time = None
self._cache = {}
logger.info("模板缓存已清除")
# 全局模板管理器实例
template_manager = TemplateManager()