Files

139 lines
5.5 KiB
Python
Raw Permalink Normal View History

2026-09-13 21:22:54 +08:00
"""幂等执行 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())