168 lines
6.5 KiB
Python
168 lines
6.5 KiB
Python
"""SQLite 数据库连接与初始化(分库)
|
||
|
||
- assets.db :Provider / Asset / VPSDetail / DomainDetail / AIAccount / Account
|
||
- metrics.db:MetricPoint / ServerInfo / SecurityCheck / EventLog
|
||
"""
|
||
|
||
from pathlib import Path
|
||
from typing import Generator
|
||
|
||
from sqlmodel import Session, SQLModel, create_engine
|
||
|
||
BASE_DIR = Path(__file__).resolve().parent.parent
|
||
DATA_DIR = BASE_DIR / "data"
|
||
DATA_DIR.mkdir(parents=True, exist_ok=True)
|
||
|
||
ASSETS_DB_URL = f"sqlite:///{DATA_DIR / 'assets.db'}"
|
||
METRICS_DB_URL = f"sqlite:///{DATA_DIR / 'metrics.db'}"
|
||
|
||
# timeout:sqlite3 内置 busy 等待(秒),随连接创建生效。
|
||
# SQLite 默认使用 NullPool(每请求新建连接),PRAGMA busy_timeout 只在单连接有效,
|
||
# 必须通过 connect_args 传递,否则并发写入(Agent 上报 + 清理任务)会报 database is locked。
|
||
assets_engine = create_engine(
|
||
ASSETS_DB_URL,
|
||
echo=False,
|
||
connect_args={"check_same_thread": False, "timeout": 30},
|
||
)
|
||
metrics_engine = create_engine(
|
||
METRICS_DB_URL,
|
||
echo=False,
|
||
connect_args={"check_same_thread": False, "timeout": 30},
|
||
)
|
||
|
||
|
||
def _enable_wal(engine) -> None:
|
||
"""启用 WAL 模式,提升 SQLite 并发读写能力(journal_mode 持久化到库文件,设置一次即可)"""
|
||
import sqlalchemy as sa
|
||
with engine.connect() as conn:
|
||
conn.execute(sa.text("PRAGMA journal_mode=WAL"))
|
||
|
||
|
||
_enable_wal(assets_engine)
|
||
_enable_wal(metrics_engine)
|
||
|
||
# 兼容旧代码:默认 engine 指向资产库
|
||
engine = assets_engine
|
||
|
||
|
||
def init_db() -> None:
|
||
"""分库建表(幂等,可重复调用)"""
|
||
from app.models.asset import Account, AIAccount, Asset, CloudflareDetail, DomainDetail, VPSDetail # noqa: F401
|
||
from app.models.monitor import ( # noqa: F401
|
||
EventLog,
|
||
MetricPoint,
|
||
SecurityCheck,
|
||
ServerInfo,
|
||
)
|
||
from app.models.provider import Provider # noqa: F401
|
||
from app.models.ssl import SiteCert, Subdomain # noqa: F401
|
||
|
||
asset_models = [Provider, Asset, VPSDetail, DomainDetail, AIAccount, CloudflareDetail, Subdomain, SiteCert, Account]
|
||
metric_models = [MetricPoint, ServerInfo, SecurityCheck, EventLog]
|
||
|
||
SQLModel.metadata.create_all(
|
||
assets_engine, tables=[m.__table__ for m in asset_models]
|
||
)
|
||
SQLModel.metadata.create_all(
|
||
metrics_engine, tables=[m.__table__ for m in metric_models]
|
||
)
|
||
_migrate_assets_db()
|
||
_migrate_indexes()
|
||
|
||
|
||
def _migrate_assets_db() -> None:
|
||
"""轻量迁移:为已有表补充新增列(SQLite create_all 不会修改已有表结构)"""
|
||
import sqlalchemy as sa
|
||
|
||
with assets_engine.begin() as conn:
|
||
insp = sa.inspect(conn)
|
||
if insp.has_table("assets"):
|
||
cols = {c["name"] for c in insp.get_columns("assets")}
|
||
if "external_id" not in cols:
|
||
conn.execute(sa.text("ALTER TABLE assets ADD COLUMN external_id VARCHAR"))
|
||
if "renew_url" not in cols:
|
||
conn.execute(sa.text("ALTER TABLE assets ADD COLUMN renew_url VARCHAR"))
|
||
if "cancel_url" not in cols:
|
||
conn.execute(sa.text("ALTER TABLE assets ADD COLUMN cancel_url VARCHAR"))
|
||
if insp.has_table("providers"):
|
||
cols = {c["name"] for c in insp.get_columns("providers")}
|
||
if "last_synced_at" not in cols:
|
||
conn.execute(sa.text("ALTER TABLE providers ADD COLUMN last_synced_at DATETIME"))
|
||
if "services" not in cols:
|
||
conn.execute(sa.text("ALTER TABLE providers ADD COLUMN services VARCHAR DEFAULT ''"))
|
||
# 为已有预设平台补充 services(仅补空值,用户自定义行不动)
|
||
_backfill_provider_services(conn)
|
||
if insp.has_table("ai_accounts"):
|
||
cols = {c["name"] for c in insp.get_columns("ai_accounts")}
|
||
if "api_key_encrypted" not in cols:
|
||
conn.execute(sa.text("ALTER TABLE ai_accounts ADD COLUMN api_key_encrypted VARCHAR"))
|
||
# 迁移:将明文 api_key 加密后存入 api_key_encrypted,并清空原字段
|
||
_migrate_plaintext_api_keys(conn)
|
||
|
||
|
||
def _migrate_indexes() -> None:
|
||
"""为已有数据库补充复合索引(create_all 不会为已存在的表补索引,IF NOT EXISTS 幂等)"""
|
||
import sqlalchemy as sa
|
||
|
||
stmts = [
|
||
(assets_engine, "CREATE INDEX IF NOT EXISTS ix_assets_provider_ext_type ON assets (provider_id, external_id, asset_type)"),
|
||
(metrics_engine, "CREATE INDEX IF NOT EXISTS ix_metric_points_asset_ts ON metric_points (asset_id, ts)"),
|
||
(metrics_engine, "CREATE INDEX IF NOT EXISTS ix_security_checks_asset_ts ON security_checks (asset_id, ts)"),
|
||
]
|
||
for eng, sql in stmts:
|
||
with eng.begin() as conn:
|
||
conn.execute(sa.text(sql))
|
||
|
||
|
||
def _migrate_plaintext_api_keys(conn) -> None:
|
||
"""一次性迁移:将 ai_accounts 中残留的明文 api_key 加密后存入 api_key_encrypted,并清空原字段"""
|
||
import sqlalchemy as sa
|
||
from app.core.crypto import encrypt
|
||
|
||
rows = conn.execute(
|
||
sa.text("SELECT id, api_key FROM ai_accounts WHERE api_key IS NOT NULL AND api_key != ''")
|
||
).fetchall()
|
||
if not rows:
|
||
return
|
||
for row in rows:
|
||
encrypted = encrypt(row[1])
|
||
if encrypted:
|
||
conn.execute(
|
||
sa.text("UPDATE ai_accounts SET api_key_encrypted = :enc, api_key = NULL WHERE id = :id"),
|
||
{"enc": encrypted, "id": row[0]},
|
||
)
|
||
|
||
|
||
def _backfill_provider_services(conn) -> None:
|
||
"""为已有预设平台补充 services 服务列表(只更新 services 为空的预设 slug)"""
|
||
import sqlalchemy as sa
|
||
from app.core.seed import PRESET_PROVIDERS
|
||
|
||
for data in PRESET_PROVIDERS:
|
||
services = data.get("services")
|
||
if not services:
|
||
continue
|
||
conn.execute(
|
||
sa.text(
|
||
"UPDATE providers SET services = :services "
|
||
"WHERE slug = :slug AND (services IS NULL OR services = '')"
|
||
),
|
||
{"services": services, "slug": data["slug"]},
|
||
)
|
||
|
||
|
||
def get_session() -> Generator[Session, None, None]:
|
||
"""资产库会话(默认)
|
||
|
||
expire_on_commit=False:commit 后不失效对象属性,避免后续访问触发隐式
|
||
重新加载查询;需要最新值的场景由调用方显式 session.refresh()。
|
||
"""
|
||
with Session(assets_engine, expire_on_commit=False) as session:
|
||
yield session
|
||
|
||
|
||
def get_metrics_session() -> Generator[Session, None, None]:
|
||
"""监控 / 日志库会话"""
|
||
with Session(metrics_engine, expire_on_commit=False) as session:
|
||
yield session
|