Files
CapacityReport/app/db_init.py
T

156 lines
5.8 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""数据库前置检查:先确保各「库」存在(不存在则自动创建),再确保各库「必须存在的结构表」已建好。
约定(见 db_init/README.md):
- 建库:对 celldata 与 capacityreport(仅直连 MySQL 时) 两个目标库执行 CREATE DATABASE IF NOT EXISTS,
库名以用户配置为准;连接时若库不存在不会报错(先用不带 database 的连接建库)。
- 建表:初始化 SQL 放在 BASE_DIR/db_init/,文件名形如 ``<库标识>.<表名>.sql``,仅当目标表「不存在」时执行,
因此 sector_band_ref 这类带预设数据的表只在首次创建时写入,绝不覆盖用户自定义。
- 库标识 -> 实际连接(库名以用户配置为准,不写死):
celldata -> AppConfig.cell_data.mysql(CellData 库,始终直连 MySQL)
capacityreport -> AppConfig.mysql(主仓库,仅当 warehouse_type == 'mysql' 直连时检查)
- 全过程 best-effort:单个库/单张表失败只记录日志并继续,不阻断启动或处理流程。
"""
from __future__ import annotations
from pathlib import Path
from typing import Optional
import pymysql
from app.config import BASE_DIR, AppConfig, MySQLConfig
DB_INIT_DIR = BASE_DIR / "db_init"
# 需要确保存在的目标库(建库不依赖是否有建表 SQL;主库表为动态生成,故无建表文件但仍需建库)
DB_KEYS = ("celldata", "capacityreport")
def _log(logger, message: str) -> None:
if logger is not None:
try:
logger.info(message)
return
except Exception:
pass
print(message)
def _target_mysql(app_config: AppConfig, db_key: str) -> Optional[MySQLConfig]:
"""库标识 -> MySQL 连接配置;不适用(如主库走 Metrix)时返回 None。"""
if db_key == "celldata":
return app_config.cell_data.mysql.normalized()
if db_key == "capacityreport":
if app_config.warehouse_type != "mysql":
return None
return app_config.mysql.normalized()
return None
def _ensure_database(mysql: MySQLConfig) -> None:
"""库不存在则创建(用不带 database 的连接执行 CREATE DATABASE IF NOT EXISTS)。"""
safe_name = str(mysql.dbname).replace("`", "")
conn = pymysql.connect(
host=mysql.host,
port=mysql.port,
user=mysql.user,
password=mysql.passwd,
charset="utf8mb4",
cursorclass=pymysql.cursors.Cursor,
autocommit=True,
)
try:
with conn.cursor() as cursor:
cursor.execute(
f"CREATE DATABASE IF NOT EXISTS `{safe_name}` "
f"CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci"
)
finally:
conn.close()
def _ensure_databases(app_config: AppConfig, logger=None) -> None:
"""确保所有目标库存在(best-effort)。"""
for db_key in DB_KEYS:
mysql = _target_mysql(app_config, db_key)
if mysql is None or not mysql.dbname:
continue
try:
_ensure_database(mysql)
except Exception as exc: # noqa: BLE001
_log(logger, f"[前置检查] 确保库 {db_key}({mysql.dbname}@{mysql.host}:{mysql.port}) 失败:{exc}")
def _connect(mysql: MySQLConfig):
return pymysql.connect(
host=mysql.host,
port=mysql.port,
user=mysql.user,
password=mysql.passwd,
database=mysql.dbname,
charset="utf8mb4",
cursorclass=pymysql.cursors.Cursor,
autocommit=False,
)
def _existing_tables(conn) -> set[str]:
with conn.cursor() as cursor:
cursor.execute("SHOW TABLES")
return {str(row[0]).lower() for row in cursor.fetchall()}
def _run_sql_file(conn, path: Path) -> None:
from app.processor import DataProcessor
statements = DataProcessor.parse_sql_script(path.read_text(encoding="utf-8"))
with conn.cursor() as cursor:
for statement in statements:
cursor.execute(statement)
conn.commit()
def _discover() -> dict[str, list[tuple[str, Path]]]:
"""收集 db_init 下的初始化 SQL,按库标识分组:{库标识: [(表名, 路径), ...]}。"""
groups: dict[str, list[tuple[str, Path]]] = {}
if not DB_INIT_DIR.exists():
return groups
for path in sorted(DB_INIT_DIR.glob("*.sql")):
stem = path.stem # 例如 celldata.sector_band_ref
if "." not in stem:
continue
db_key, table = stem.split(".", 1)
groups.setdefault(db_key, []).append((table, path))
return groups
def ensure_required_tables(app_config: AppConfig, logger=None) -> None:
"""先确保各库存在(不存在则创建),再检查各库必须存在的表,缺则按初始化 SQL 建好。失败不抛出。"""
_ensure_databases(app_config, logger)
groups = _discover()
for db_key, items in groups.items():
mysql = _target_mysql(app_config, db_key)
if mysql is None or not mysql.dbname:
continue
try:
conn = _connect(mysql)
except Exception as exc: # noqa: BLE001
_log(logger, f"[前置检查] 连接库 {db_key}({mysql.dbname}@{mysql.host}:{mysql.port}) 失败,跳过:{exc}")
continue
try:
existing = _existing_tables(conn)
for table, path in items:
if table.lower() in existing:
continue
try:
_run_sql_file(conn, path)
_log(logger, f"[前置检查] {db_key}.{table} 不存在,已按 {path.name} 初始化建表")
except Exception as exc: # noqa: BLE001
try:
conn.rollback()
except Exception:
pass
_log(logger, f"[前置检查] 初始化 {db_key}.{table} 失败:{exc}")
finally:
conn.close()