Files
Mutual_Fund/scripts/apply_memory_unit_upgrade.py

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