baodan/api/insurance/db/migrate_019.py

162 lines
6.3 KiB
Python
Raw Normal View History

"""迁移 019: 修复已有 schema 和数据问题。
修复内容
1. 转换 slides_config_json 中的对象格式 {"slides": [...]} 为数组格式 [...]
2. poster_records 表添加生成方式追踪字段
3. 为常用查询添加索引
"""
import json
import logging
from sqlalchemy import text, inspect
logger = logging.getLogger(__name__)
def migrate():
"""执行迁移。"""
from insurance.db.compat import db
2026-07-27 15:40:58 +08:00
dialect = db.engine.dialect.name
_fix_slides_config_json(db)
_add_poster_record_fields(db, dialect)
_add_indexes(db, dialect)
db.session.commit()
def _column_exists(db, table_name: str, column_name: str) -> bool:
"""检查列是否存在。"""
2026-07-27 15:40:58 +08:00
dialect = db.engine.dialect.name
if dialect == "postgresql":
result = db.session.execute(text(
"SELECT COUNT(*) FROM information_schema.columns "
"WHERE table_name = :table AND column_name = :column"
), {"table": table_name, "column": column_name})
return result.scalar() > 0
else:
try:
result = db.session.execute(text(f"PRAGMA table_info({table_name})"))
return any(row[1] == column_name for row in result.fetchall())
except Exception:
return False
def _fix_slides_config_json(db):
"""将对象格式 {"slides": [...]} 转换为数组格式 [...]。
历史数据中 migrate_017 写入的是 {"slides": [...]}
但渲染器和前端均期望数组格式
"""
templates = db.session.execute(text(
"SELECT id, slides_config_json FROM insurance_ppt_templates "
"WHERE slides_config_json IS NOT NULL"
)).fetchall()
fixed_count = 0
for tpl_id, config_json in templates:
if not config_json or not config_json.strip():
continue
try:
config = json.loads(config_json)
except (json.JSONDecodeError, TypeError):
logger.warning(f"模板 {tpl_id} 的 slides_config_json 无法解析,跳过")
continue
# 如果是对象格式 {"slides": [...]},转换为数组
if isinstance(config, dict) and "slides" in config:
slides_array = config["slides"]
if isinstance(slides_array, list):
db.session.execute(text(
"UPDATE insurance_ppt_templates SET slides_config_json = :config WHERE id = :id"
), {"config": json.dumps(slides_array, ensure_ascii=False), "id": tpl_id})
fixed_count += 1
logger.info(f"模板 {tpl_id}: 对象格式已转换为数组格式")
else:
logger.warning(f"模板 {tpl_id}: slides 值不是数组,跳过")
elif isinstance(config, list):
pass # 已经是数组格式,无需处理
else:
logger.warning(f"模板 {tpl_id}: 未知格式 {type(config).__name__},跳过")
if fixed_count > 0:
logger.info(f"已修复 {fixed_count} 个模板的 slides_config_json 格式")
def _add_poster_record_fields(db, dialect):
"""为 poster_records 表添加生成方式追踪和任务状态字段。"""
new_columns = [
("generation_mode", "VARCHAR(20)", "生成方式: ai/fallback"),
("image_provider", "VARCHAR(50)", "图片供应商"),
("image_model", "VARCHAR(100)", "图片模型名称"),
("task_status", "VARCHAR(20)", "任务状态: pending/queued/generating/done/failed"),
("task_progress", "INTEGER", "任务进度 0-100"),
("task_error", "TEXT", "任务错误信息"),
("started_at", "TIMESTAMP", "任务开始时间"),
("finished_at", "TIMESTAMP", "任务完成时间"),
]
for col_name, col_type, comment in new_columns:
if _column_exists(db, "poster_records", col_name):
logger.debug(f"poster_records.{col_name} 已存在,跳过")
continue
try:
db.session.execute(text(
f"ALTER TABLE poster_records ADD COLUMN {col_name} {col_type}"
))
db.session.commit()
logger.info(f"已添加 poster_records.{col_name} ({comment})")
except Exception as e:
db.session.rollback()
logger.warning(f"添加 poster_records.{col_name} 失败(已回滚,不影响其他字段): {e}")
def _add_indexes(db, dialect):
"""为常用查询添加索引。"""
indexes = [
("idx_poster_records_user_id", "poster_records", ["user_id"]),
("idx_poster_records_created_at", "poster_records", ["created_at"]),
("idx_poster_case_uploads_user_id", "poster_case_uploads", ["user_id"]),
("idx_ppt_history_user_id", "insurance_ppt_history", ["user_id"]),
("idx_ppt_history_created_at", "insurance_ppt_history", ["created_at"]),
("idx_system_settings_key", "system_settings", ["key"]),
]
for idx_name, table_name, columns in indexes:
try:
# 检查表是否存在
if dialect == "postgresql":
result = db.session.execute(text(
"SELECT COUNT(*) FROM information_schema.tables WHERE table_name = :table"
), {"table": table_name})
if result.scalar() == 0:
continue
else:
try:
db.session.execute(text(f"SELECT 1 FROM {table_name} LIMIT 0"))
except Exception:
continue
# 检查索引是否已存在
if dialect == "postgresql":
result = db.session.execute(text(
"SELECT COUNT(*) FROM pg_indexes WHERE indexname = :idx"
), {"idx": idx_name})
if result.scalar() > 0:
continue
else:
try:
result = db.session.execute(text(
"SELECT COUNT(*) FROM sqlite_master WHERE type='index' AND name = :idx"
), {"idx": idx_name})
if result.scalar() > 0:
continue
except Exception:
continue
cols = ", ".join(columns)
db.session.execute(text(f"CREATE INDEX {idx_name} ON {table_name} ({cols})"))
logger.info(f"已创建索引 {idx_name}")
except Exception as e:
logger.warning(f"创建索引 {idx_name} 失败: {e}")
db.session.rollback()