""" 数据库配置和连接管理 """ 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