Files
taiji-AI-PAD/services/mcp-server/database.py
T
2025-12-25 07:56:49 +00:00

456 lines
14 KiB
Python

"""
数据库配置和连接管理
"""
import asyncio
import ssl
from typing import AsyncGenerator
from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine, async_sessionmaker
from sqlalchemy.orm import sessionmaker
from sqlalchemy import text
from sqlalchemy.engine import make_url
import logging
from config import settings
from models import Base
logger = logging.getLogger(__name__)
database_url = settings.database_url
database_url_obj = make_url(database_url)
engine_kwargs = {
"echo": settings.debug,
"pool_pre_ping": True,
}
if database_url_obj.get_backend_name().startswith("sqlite"):
engine_kwargs["connect_args"] = {"check_same_thread": False}
else:
engine_kwargs.update({
"pool_size": 20,
"max_overflow": 0,
"pool_recycle": 3600,
})
if database_url_obj.host and database_url_obj.host.endswith("postgres.database.azure.com"):
# Azure Database for PostgreSQL requires TLS; provide a default SSL context.
ssl_context = ssl.create_default_context()
existing_connect_args = engine_kwargs.get("connect_args") or {}
existing_connect_args["ssl"] = ssl_context
engine_kwargs["connect_args"] = existing_connect_args
# 创建异步数据库引擎
engine = create_async_engine(database_url, **engine_kwargs)
# 判断当前是否使用SQLite后端
def is_sqlite_backend() -> bool:
backend = engine.url.get_backend_name()
return backend.startswith("sqlite")
# 创建异步会话工厂
AsyncSessionLocal = async_sessionmaker(
engine,
class_=AsyncSession,
expire_on_commit=False
)
async def get_db() -> AsyncGenerator[AsyncSession, None]:
"""获取数据库会话的依赖注入函数"""
async with AsyncSessionLocal() as session:
try:
yield session
except Exception as e:
logger.error(f"数据库会话错误: {e}")
await session.rollback()
raise
finally:
await session.close()
async def init_db():
"""初始化数据库"""
try:
# 创建所有表
async with engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
logger.info("数据库表创建成功")
# 尝试创建初始数据,失败不阻止启动
try:
await create_initial_data()
except Exception as init_err:
logger.warning(f"创建初始数据失败(服务将继续启动): {init_err}")
except Exception as e:
logger.error(f"数据库初始化失败: {e}")
raise
async def create_initial_data():
"""创建初始数据"""
try:
async with AsyncSessionLocal() as session:
# 检查是否已有数据
result = await session.execute(text("SELECT COUNT(*) FROM users"))
user_count = result.scalar()
if user_count == 0:
# 创建默认管理员用户
from models import User
from passlib.context import CryptContext
pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
admin_user = User(
username="admin",
email="admin@taiji-ai.com",
hashed_password=pwd_context.hash("admin123"),
full_name="系统管理员",
is_active=True,
is_admin=True
)
session.add(admin_user)
await session.commit()
logger.info("默认管理员用户创建成功")
# 创建示例工具
await create_sample_tools(session)
except Exception as e:
logger.error(f"创建初始数据失败: {e}")
raise
async def create_sample_tools(session: AsyncSession):
"""创建示例工具"""
try:
from models import Tool
# 检查是否已有工具
result = await session.execute(text("SELECT COUNT(*) FROM tools"))
tool_count = result.scalar()
if tool_count == 0:
# 创建示例工具
sample_tools = [
{
"name": "web_search",
"description": "网络搜索工具",
"category": "api",
"schema": {
"type": "object",
"properties": {
"query": {
"type": "string",
"description": "搜索查询"
},
"limit": {
"type": "integer",
"description": "结果数量限制",
"default": 10
}
},
"required": ["query"]
},
"endpoint": "https://api.example.com/search",
"method": "POST",
"auth_type": "api_key",
"rate_limit": 100,
"cost_per_call": 0.01,
"is_public": True
},
{
"name": "text_completion",
"description": "文本补全工具",
"category": "llm",
"schema": {
"type": "object",
"properties": {
"prompt": {
"type": "string",
"description": "输入提示"
},
"max_tokens": {
"type": "integer",
"description": "最大token数",
"default": 150
},
"temperature": {
"type": "number",
"description": "温度参数",
"default": 0.7
}
},
"required": ["prompt"]
},
"rate_limit": 60,
"cost_per_call": 0.05,
"is_public": True
},
{
"name": "weather_api",
"description": "天气查询API",
"category": "api",
"schema": {
"type": "object",
"properties": {
"city": {
"type": "string",
"description": "城市名称"
},
"unit": {
"type": "string",
"enum": ["celsius", "fahrenheit"],
"default": "celsius"
}
},
"required": ["city"]
},
"endpoint": "https://api.openweathermap.org/data/2.5/weather",
"method": "GET",
"auth_type": "api_key",
"rate_limit": 1000,
"cost_per_call": 0.001,
"is_public": True
}
]
for tool_data in sample_tools:
tool = Tool(**tool_data)
session.add(tool)
await session.commit()
logger.info(f"创建了 {len(sample_tools)} 个示例工具")
except Exception as e:
logger.error(f"创建示例工具失败: {e}")
raise
async def check_db_connection():
"""检查数据库连接"""
try:
async with AsyncSessionLocal() as session:
await session.execute(text("SELECT 1"))
return True
except Exception as e:
logger.error(f"数据库连接检查失败: {e}")
return False
async def get_db_stats():
"""获取数据库统计信息"""
try:
async with AsyncSessionLocal() as session:
stats = {}
# 获取各表的记录数
tables = ["users", "agents", "tools", "sessions", "executions", "billing"]
for table in tables:
result = await session.execute(text(f"SELECT COUNT(*) FROM {table}"))
stats[table] = result.scalar()
return stats
except Exception as e:
logger.error(f"获取数据库统计失败: {e}")
return {}
async def cleanup_old_records():
"""清理旧记录"""
try:
async with AsyncSessionLocal() as session:
# 清理超过30天的执行记录
if is_sqlite_backend():
result = await session.execute(text("""
DELETE FROM executions
WHERE created_at < datetime('now', '-30 days')
"""))
else:
result = await session.execute(text("""
DELETE FROM executions
WHERE created_at < NOW() - INTERVAL '30 days'
"""))
deleted_executions = result.rowcount
# 清理超过7天的会话记录
if is_sqlite_backend():
result = await session.execute(text("""
DELETE FROM sessions
WHERE created_at < datetime('now', '-7 days')
AND status != 'active'
"""))
else:
result = await session.execute(text("""
DELETE FROM sessions
WHERE created_at < NOW() - INTERVAL '7 days'
AND status != 'active'
"""))
deleted_sessions = result.rowcount
await session.commit()
logger.info(f"清理完成: 删除了 {deleted_executions} 条执行记录, {deleted_sessions} 条会话记录")
return {
"deleted_executions": deleted_executions,
"deleted_sessions": deleted_sessions
}
except Exception as e:
logger.error(f"清理旧记录失败: {e}")
return {}
async def backup_db():
"""数据库备份"""
try:
import subprocess
from datetime import datetime
import os
import shutil
# 生成备份文件名
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
backup_dir = "/app/backups"
os.makedirs(backup_dir, exist_ok=True)
if is_sqlite_backend():
db_path = engine.url.database
if not db_path:
raise ValueError("SQLite数据库路径为空,无法备份")
db_path = os.path.abspath(db_path)
backup_file = os.path.join(backup_dir, f"mcp_sqlite_backup_{timestamp}.db")
shutil.copy2(db_path, backup_file)
logger.info(f"SQLite数据库备份成功: {backup_file}")
return backup_file
backup_file = f"{backup_dir}/taiji_db_backup_{timestamp}.sql"
# 执行pg_dump命令
cmd = [
"pg_dump",
settings.database_url.replace("postgresql+asyncpg://", "postgresql://"),
"-f", backup_file
]
result = subprocess.run(cmd, capture_output=True, text=True)
if result.returncode == 0:
logger.info(f"数据库备份成功: {backup_file}")
return backup_file
else:
logger.error(f"数据库备份失败: {result.stderr}")
return None
except Exception as e:
logger.error(f"数据库备份异常: {e}")
return None
async def close_db():
"""关闭数据库连接"""
try:
await engine.dispose()
logger.info("数据库连接已关闭")
except Exception as e:
logger.error(f"关闭数据库连接失败: {e}")
# 数据库事件处理
async def on_startup():
"""应用启动时的数据库操作"""
await init_db()
async def on_shutdown():
"""应用关闭时的数据库操作"""
await close_db()
# 定期清理任务
async def periodic_cleanup():
"""定期清理任务"""
while True:
try:
await asyncio.sleep(3600) # 每小时执行一次
await cleanup_old_records()
except Exception as e:
logger.error(f"定期清理任务异常: {e}")
# 数据库迁移辅助函数
async def migrate_db():
"""数据库迁移(简化版本)"""
try:
# 这里可以添加数据迁移逻辑
# 在生产环境中应该使用Alembic进行数据库版本管理
logger.info("数据库迁移检查完成")
except Exception as e:
logger.error(f"数据库迁移失败: {e}")
raise
# 性能优化
async def optimize_db():
"""数据库性能优化"""
try:
async with AsyncSessionLocal() as session:
# 更新表统计信息
await session.execute(text("ANALYZE;"))
# 重建索引(如果需要)
# await session.execute(text("REINDEX DATABASE taiji_db;"))
await session.commit()
logger.info("数据库优化完成")
except Exception as e:
logger.error(f"数据库优化失败: {e}")
# 健康检查
async def health_check() -> dict:
"""数据库健康检查"""
health_info = {
"database": "unknown",
"connection_pool": "unknown",
"stats": {}
}
try:
# 检查连接
if await check_db_connection():
health_info["database"] = "healthy"
else:
health_info["database"] = "unhealthy"
# 检查连接池状态
pool = engine.pool
health_info["connection_pool"] = {
"size": pool.size(),
"checked_in": pool.checkedin(),
"checked_out": pool.checkedout()
}
# 获取统计信息
health_info["stats"] = await get_db_stats()
except Exception as e:
logger.error(f"数据库健康检查失败: {e}")
health_info["database"] = "error"
health_info["error"] = str(e)
return health_info