forked from xiaohei/taiji-AI-PAD
456 lines
14 KiB
Python
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
|
|
|