forked from zhanggangyong/agent_management
161 lines
6.4 KiB
Python
161 lines
6.4 KiB
Python
"""
|
|
数据库迁移脚本 - 添加多框架支持字段
|
|
运行: python migrate_multi_framework.py
|
|
"""
|
|
import sys
|
|
import os
|
|
|
|
# 添加父目录到路径
|
|
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
|
|
|
from database import engine, Base
|
|
from sqlalchemy import text
|
|
import logging
|
|
|
|
logging.basicConfig(level=logging.INFO)
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
def migrate():
|
|
"""执行数据库迁移"""
|
|
|
|
logger.info("🚀 开始数据库迁移 - 添加多框架支持")
|
|
|
|
with engine.connect() as conn:
|
|
# 开始事务
|
|
trans = conn.begin()
|
|
|
|
try:
|
|
# 检查数据库类型
|
|
db_url = str(engine.url)
|
|
is_sqlite = 'sqlite' in db_url
|
|
is_postgres = 'postgres' in db_url
|
|
|
|
logger.info(f"数据库类型: {'SQLite' if is_sqlite else 'PostgreSQL' if is_postgres else 'Unknown'}")
|
|
|
|
# Templates 表迁移
|
|
logger.info("📝 迁移 templates 表...")
|
|
|
|
migrations_templates = [
|
|
("agent_framework", "VARCHAR(50)", "langchain"),
|
|
("tools_config", "JSON" if is_postgres else "TEXT", None),
|
|
("default_model_provider", "VARCHAR(100)", None),
|
|
("default_model_name", "VARCHAR(200)", None),
|
|
]
|
|
|
|
for column_name, column_type, default_value in migrations_templates:
|
|
try:
|
|
if default_value:
|
|
if is_sqlite:
|
|
# SQLite 需要特殊处理
|
|
conn.execute(text(f"ALTER TABLE templates ADD COLUMN {column_name} {column_type} DEFAULT '{default_value}'"))
|
|
else:
|
|
conn.execute(text(f"ALTER TABLE templates ADD COLUMN {column_name} {column_type} DEFAULT '{default_value}'"))
|
|
else:
|
|
conn.execute(text(f"ALTER TABLE templates ADD COLUMN {column_name} {column_type}"))
|
|
logger.info(f" ✅ 添加列: templates.{column_name}")
|
|
except Exception as e:
|
|
if "already exists" in str(e) or "duplicate column" in str(e).lower():
|
|
logger.info(f" ⏭️ 跳过已存在的列: templates.{column_name}")
|
|
else:
|
|
raise
|
|
|
|
# Agents 表迁移
|
|
logger.info("📝 迁移 agents 表...")
|
|
|
|
migrations_agents = [
|
|
("agent_framework", "VARCHAR(50)", "langchain"),
|
|
("tools_config", "JSON" if is_postgres else "TEXT", None),
|
|
("tool_endpoint", "VARCHAR(500)", None),
|
|
("tool_api_key", "VARCHAR(500)", None),
|
|
("model_provider", "VARCHAR(100)", None),
|
|
("model_name", "VARCHAR(200)", None),
|
|
("model_endpoint", "VARCHAR(500)", None),
|
|
("model_api_key", "VARCHAR(500)", None),
|
|
("storage_connection_string", "VARCHAR(1000)", None),
|
|
("storage_account_name", "VARCHAR(200)", None),
|
|
]
|
|
|
|
for column_name, column_type, default_value in migrations_agents:
|
|
try:
|
|
if default_value:
|
|
if is_sqlite:
|
|
conn.execute(text(f"ALTER TABLE agents ADD COLUMN {column_name} {column_type} DEFAULT '{default_value}'"))
|
|
else:
|
|
conn.execute(text(f"ALTER TABLE agents ADD COLUMN {column_name} {column_type} DEFAULT '{default_value}'"))
|
|
else:
|
|
conn.execute(text(f"ALTER TABLE agents ADD COLUMN {column_name} {column_type}"))
|
|
logger.info(f" ✅ 添加列: agents.{column_name}")
|
|
except Exception as e:
|
|
if "already exists" in str(e) or "duplicate column" in str(e).lower():
|
|
logger.info(f" ⏭️ 跳过已存在的列: agents.{column_name}")
|
|
else:
|
|
raise
|
|
|
|
# 提交事务
|
|
trans.commit()
|
|
logger.info("✅ 数据库迁移完成")
|
|
|
|
except Exception as e:
|
|
# 回滚事务
|
|
trans.rollback()
|
|
logger.error(f"❌ 迁移失败: {str(e)}")
|
|
raise
|
|
|
|
|
|
def verify_migration():
|
|
"""验证迁移结果"""
|
|
logger.info("🔍 验证迁移结果...")
|
|
|
|
with engine.connect() as conn:
|
|
# 检查 templates 表
|
|
result = conn.execute(text("SELECT * FROM templates LIMIT 0"))
|
|
templates_columns = result.keys()
|
|
logger.info(f"Templates 表列: {list(templates_columns)}")
|
|
|
|
# 检查 agents 表
|
|
result = conn.execute(text("SELECT * FROM agents LIMIT 0"))
|
|
agents_columns = result.keys()
|
|
logger.info(f"Agents 表列: {list(agents_columns)}")
|
|
|
|
# 验证新字段
|
|
required_template_columns = [
|
|
'agent_framework', 'tools_config',
|
|
'default_model_provider', 'default_model_name'
|
|
]
|
|
|
|
required_agent_columns = [
|
|
'agent_framework', 'tools_config', 'tool_endpoint', 'tool_api_key',
|
|
'model_provider', 'model_name', 'model_endpoint', 'model_api_key',
|
|
'storage_connection_string', 'storage_account_name'
|
|
]
|
|
|
|
missing_template_cols = [col for col in required_template_columns if col not in templates_columns]
|
|
missing_agent_cols = [col for col in required_agent_columns if col not in agents_columns]
|
|
|
|
if missing_template_cols:
|
|
logger.warning(f"⚠️ Templates 表缺少列: {missing_template_cols}")
|
|
else:
|
|
logger.info("✅ Templates 表所有必需列都存在")
|
|
|
|
if missing_agent_cols:
|
|
logger.warning(f"⚠️ Agents 表缺少列: {missing_agent_cols}")
|
|
else:
|
|
logger.info("✅ Agents 表所有必需列都存在")
|
|
|
|
if not missing_template_cols and not missing_agent_cols:
|
|
logger.info("🎉 迁移验证成功!")
|
|
return True
|
|
else:
|
|
logger.error("❌ 迁移验证失败")
|
|
return False
|
|
|
|
|
|
if __name__ == "__main__":
|
|
try:
|
|
migrate()
|
|
verify_migration()
|
|
except Exception as e:
|
|
logger.error(f"迁移过程出错: {str(e)}")
|
|
sys.exit(1)
|