"""幂等执行 memory_unit 表结构升级(对齐 model/memory_unit.py)。 用法: python scripts/apply_memory_unit_upgrade.py 对应 SQL 版本见 sql/memory_unit_upgrade_20260913.sql。 重复执行安全:ADD COLUMN 前检查 information_schema,MODIFY/UPDATE 本身幂等。 """ from __future__ import annotations import asyncio import sys from pathlib import Path sys.path.insert(0, str(Path(__file__).resolve().parents[1])) from sqlalchemy import text from config.database import mysql as mysql_db from config.database.mysql import get_session_factory from config.settings import settings TABLE = "memory_unit" ADD_COLUMNS = { "session_id": "session_id VARCHAR(64) NULL COMMENT '产生记忆的会话ID'", "agent_run_id": "agent_run_id VARCHAR(64) NULL COMMENT '产生记忆的Agent运行ID'", "evidence_ref": "evidence_ref JSON NULL COMMENT '证据引用列表(会话/消息溯源)'", "historical_accuracy": "historical_accuracy DECIMAL(5,2) NOT NULL DEFAULT 0.50 COMMENT '历史准确率'", "confidence_version": "confidence_version VARCHAR(32) NULL COMMENT '置信度算法版本'", "confidence_reason": "confidence_reason VARCHAR(255) NULL COMMENT '置信度评分原因'", "confidence_update_time": "confidence_update_time DATETIME NULL COMMENT '置信度更新时间'", "memory_version": "memory_version INT NOT NULL DEFAULT 1 COMMENT '记忆版本号'", "last_verified_at": "last_verified_at DATETIME NULL COMMENT '最近验证时间'", "milvus_id": "milvus_id VARCHAR(128) NULL COMMENT 'Milvus向量主键'", "graph_node_id": "graph_node_id VARCHAR(128) NULL COMMENT 'Neo4j图谱节点ID'", "milvus_sync_status": "milvus_sync_status VARCHAR(16) NOT NULL DEFAULT 'pending' COMMENT '向量同步状态'", "neo4j_sync_status": "neo4j_sync_status VARCHAR(16) NOT NULL DEFAULT 'pending' COMMENT '图谱同步状态'", "sync_retry_count": "sync_retry_count INT NOT NULL DEFAULT 0 COMMENT '同步重试次数'", "last_sync_error": "last_sync_error VARCHAR(500) NULL COMMENT '最近同步错误'", "next_retry_at": "next_retry_at DATETIME NULL COMMENT '下次重试时间'", "last_synced_at": "last_synced_at DATETIME NULL COMMENT '最近成功同步时间'", } async def columns_of(session, database: str) -> dict[str, str]: rows = ( await session.execute( text( "SELECT COLUMN_NAME, DATA_TYPE FROM information_schema.columns " "WHERE TABLE_SCHEMA = :d AND TABLE_NAME = :t" ), {"d": database, "t": TABLE}, ) ).mappings().all() return {str(r["COLUMN_NAME"]): str(r["DATA_TYPE"]).lower() for r in rows} async def main() -> None: async with get_session_factory()() as session: db = settings.mysql.database cols = await columns_of(session, db) added = [] for name, ddl in ADD_COLUMNS.items(): if name in cols: print(f"[skip] 列已存在: {name}") continue await session.execute(text(f"ALTER TABLE {TABLE} ADD COLUMN {ddl}")) added.append(name) print(f"[ok] 新增列: {name}") await session.commit() cols = await columns_of(session, db) if cols.get("valid_from") == "date" or cols.get("valid_until") == "date": await session.execute( text( f"ALTER TABLE {TABLE} " "MODIFY COLUMN valid_from DATETIME NULL COMMENT '生效起始时间', " "MODIFY COLUMN valid_until DATETIME NULL COMMENT '失效时间'" ) ) await session.commit() print("[ok] valid_from/valid_until: DATE -> DATETIME") else: print("[skip] valid_from/valid_until 已是 DATETIME") result = await session.execute( text(f"UPDATE {TABLE} SET status = 'candidate' WHERE status = 'active'") ) await session.commit() print(f"[ok] status active->candidate: {result.rowcount} 行") await session.execute( text( f"ALTER TABLE {TABLE} MODIFY COLUMN status VARCHAR(16) NOT NULL " "DEFAULT 'candidate' COMMENT '记忆状态(candidate/confirmed/expired/rejected/archived)'" ) ) await session.commit() print("[ok] status 默认值改为 candidate") result = await session.execute( text( f"UPDATE {TABLE} SET memory_type = 'SERVICE_FACT' " "WHERE memory_type IS NULL OR memory_type = 'FACT'" ) ) await session.commit() print(f"[ok] memory_type NULL/FACT -> SERVICE_FACT: {result.rowcount} 行") await session.execute( text( f"ALTER TABLE {TABLE} MODIFY COLUMN memory_type VARCHAR(32) NOT NULL " "COMMENT '记忆业务类型'" ) ) await session.commit() print("[ok] memory_type 收紧为 NOT NULL") final = await columns_of(session, db) missing = [name for name in ADD_COLUMNS if name not in final] if missing: print(f"[fail] 仍有缺失列: {missing}") raise SystemExit(1) null_types = ( await session.execute( text(f"SELECT COUNT(*) AS n FROM {TABLE} WHERE memory_type IS NULL") ) ).scalar() print(f"[done] 迁移完成,共新增 {len(added)} 列;memory_type NULL 残留: {null_types}") await mysql_db.dispose() if __name__ == "__main__": asyncio.run(main())