139 lines
5.5 KiB
Python
139 lines
5.5 KiB
Python
"""幂等执行 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())
|