from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker from sqlalchemy.orm import DeclarativeBase from backend.config import get_settings settings = get_settings() engine = create_async_engine(settings.database_url, echo=False) async_session = async_sessionmaker(engine, class_=AsyncSession, expire_on_commit=False) class Base(DeclarativeBase): pass async def get_db(): async with async_session() as session: try: yield session await session.commit() except Exception: await session.rollback() raise # ────────────── CRUD 工具函数 ────────────── async def db_add(db: AsyncSession, obj, *, refresh: bool = False): """创建记录并提交。refresh=True 会在提交后刷新对象(获取生成的 id 等)。""" db.add(obj) await db.flush() if refresh: await db.refresh(obj) await db.commit() async def db_update(db: AsyncSession, obj, data: dict): """批量更新字段并提交。data 一般来自 model_dump(exclude_unset=True)。""" for key, value in data.items(): setattr(obj, key, value) await db.commit() async def db_delete(db: AsyncSession, obj): """删除记录并提交。""" await db.delete(obj) await db.commit() async def db_flush(db: AsyncSession): """将待写操作刷到数据库(不提交事务),用于需要获取生成 ID 等中间状态的场景。""" await db.flush() async def init_db(): async with engine.begin() as conn: # 迁移旧表:acme_config → acme_configs try: result = await conn.execute( __import__("sqlalchemy").text("SELECT name FROM sqlite_master WHERE type='table' AND name='acme_config'") ) if result.fetchone(): await conn.execute(__import__("sqlalchemy").text("ALTER TABLE acme_config RENAME TO acme_configs")) except Exception: pass await conn.run_sync(Base.metadata.create_all) # 给 acme_configs 表补全新列(旧表迁移后可能缺少这些列) for col_def in [ "ALTER TABLE acme_configs ADD COLUMN name VARCHAR(100) NOT NULL DEFAULT '默认配置'", "ALTER TABLE acme_configs ADD COLUMN created_at DATETIME", "ALTER TABLE acme_configs ADD COLUMN updated_at DATETIME", ]: try: await conn.execute(__import__("sqlalchemy").text(col_def)) except Exception: pass # 列已存在 # 迁移:给 domains 表加 acme_config_id 列 try: await conn.execute( __import__("sqlalchemy").text("ALTER TABLE domains ADD COLUMN acme_config_id INTEGER REFERENCES acme_configs(id)") ) except Exception: pass # 列已存在 # 迁移:给 deploy_logs 表加 hostname 列 try: await conn.execute(__import__("sqlalchemy").text("ALTER TABLE deploy_logs ADD COLUMN hostname VARCHAR(200)")) except Exception: pass # 迁移:给 acme_logs 表加证书信息列 for col_def in [ "ALTER TABLE acme_logs ADD COLUMN cert_not_after DATETIME", "ALTER TABLE acme_logs ADD COLUMN cert_serial VARCHAR(100)", "ALTER TABLE acme_logs ADD COLUMN cert_san VARCHAR(500)", ]: try: await conn.execute(__import__("sqlalchemy").text(col_def)) except Exception: pass # 列已存在 # 确保至少有一条默认配置 result = await conn.execute( __import__("sqlalchemy").text("SELECT COUNT(*) FROM acme_configs") ) count = result.scalar() if count == 0: await conn.execute( __import__("sqlalchemy").text( "INSERT INTO acme_configs (name, acme_server, email, dns_provider, dns_credentials, renew_days) " "VALUES ('默认配置', 'https://acme-v02.api.letsencrypt.org/directory', '', 'aliyun', '{}', 30)" ) )