ForcePilot/scripts/migrate_kb_metadata_to_db.py

340 lines
14 KiB
Python
Raw Normal View History

import argparse
import asyncio
import glob
import json
import os
import sys
from datetime import datetime, UTC
from typing import Any
sys.path.insert(0, os.path.dirname(os.path.dirname(__file__)))
os.environ.setdefault("YUXI_SKIP_APP_INIT", "1")
from src import config
from src.repositories.evaluation_repository import EvaluationRepository
from src.repositories.knowledge_base_repository import KnowledgeBaseRepository
from src.repositories.knowledge_file_repository import KnowledgeFileRepository
from src.repositories.task_repository import TaskRepository
from src.utils import logger
def _load_json(path: str) -> dict[str, Any]:
if not os.path.exists(path):
return {}
with open(path, encoding="utf-8") as f:
return json.load(f)
def _utc_dt(value: Any) -> datetime | None:
"""Convert various datetime formats to naive UTC datetime (consistent with model)."""
if not value:
return None
if isinstance(value, datetime):
# 转换为 UTC 并移除时区信息(模型使用 DateTime 无时区)
if value.tzinfo is None:
return value
return value.astimezone(UTC).replace(tzinfo=None)
if isinstance(value, (int, float)):
# 时间戳转换为 UTC 时间
return datetime.fromtimestamp(value, tz=UTC).replace(tzinfo=None)
if isinstance(value, str):
v = value.strip()
if not v:
return None
try:
# 解析 ISO 格式并转换为 UTC
dt_val = datetime.fromisoformat(v.replace("Z", "+00:00"))
if dt_val.tzinfo is None:
return dt_val
return dt_val.astimezone(UTC).replace(tzinfo=None)
except ValueError:
return None
return None
def _default_share_config(meta: dict[str, Any]) -> dict[str, Any]:
share_config = meta.get("share_config") or {}
if "is_shared" not in share_config:
share_config["is_shared"] = True
if "accessible_departments" not in share_config:
share_config["accessible_departments"] = []
return share_config
async def rollback_all() -> None:
eval_repo = EvaluationRepository()
kb_repo = KnowledgeBaseRepository()
file_repo = KnowledgeFileRepository()
task_repo = TaskRepository()
await task_repo.delete_all()
await eval_repo.delete_all()
rows = await kb_repo.get_all()
for row in rows:
await file_repo.delete_by_db_id(row.db_id)
await kb_repo.delete(row.db_id)
async def migrate(dry_run: bool, execute: bool, rollback: bool) -> None:
from src.storage.postgres.manager import pg_manager
base_dir = os.path.join(config.save_dir, "knowledge_base_data")
global_meta_path = os.path.join(base_dir, "global_metadata.json")
global_meta = _load_json(global_meta_path).get("databases", {})
if rollback:
if dry_run:
logger.info("Dry-run rollback: would delete all knowledge metadata tables")
return
await rollback_all()
logger.info("Rollback completed")
return
# 初始化表结构
pg_manager.initialize()
await pg_manager.create_tables()
logger.info("知识库表结构初始化完成")
kb_repo = KnowledgeBaseRepository()
file_repo = KnowledgeFileRepository()
eval_repo = EvaluationRepository()
task_repo = TaskRepository()
kb_rows: list[dict[str, Any]] = []
file_rows: list[tuple[str, dict[str, Any]]] = []
benchmark_rows: list[dict[str, Any]] = []
result_rows: list[dict[str, Any]] = []
result_detail_rows: list[tuple[str, int, dict[str, Any]]] = []
kb_type_dirs = [
p for p in glob.glob(os.path.join(base_dir, "*_data")) if os.path.isdir(p) and os.path.basename(p) != "uploads"
]
for kb_dir in kb_type_dirs:
kb_type = os.path.basename(kb_dir)[: -len("_data")]
meta_file = os.path.join(kb_dir, f"metadata_{kb_type}.json")
meta = _load_json(meta_file)
databases_meta: dict[str, Any] = meta.get("databases", {})
files_meta: dict[str, Any] = meta.get("files", {})
benchmarks_meta: dict[str, Any] = meta.get("benchmarks", {})
for db_id, db_meta in databases_meta.items():
g = global_meta.get(db_id, {})
created_at = _utc_dt(g.get("created_at") or db_meta.get("created_at"))
updated_at = _utc_dt(g.get("updated_at")) or created_at
kb_rows.append(
{
"db_id": db_id,
"name": g.get("name") or db_meta.get("name") or db_id,
"description": g.get("description") or db_meta.get("description"),
"kb_type": g.get("kb_type") or db_meta.get("kb_type") or kb_type,
"embed_info": db_meta.get("embed_info") or g.get("embed_info"),
"llm_info": db_meta.get("llm_info") or g.get("llm_info"),
"query_params": db_meta.get("query_params") or g.get("query_params"),
"additional_params": g.get("additional_params") or db_meta.get("metadata") or {},
"share_config": _default_share_config(g or {}),
"mindmap": g.get("mindmap"),
"sample_questions": g.get("sample_questions") or [],
"created_at": created_at,
"updated_at": updated_at,
}
)
for file_id, fmeta in files_meta.items():
db_id = fmeta.get("database_id")
if not db_id:
continue
file_rows.append(
(
file_id,
{
"db_id": db_id,
"parent_id": fmeta.get("parent_id"),
"filename": fmeta.get("filename") or "",
"original_filename": fmeta.get("original_filename") or fmeta.get("file_name"),
"file_type": fmeta.get("file_type") or fmeta.get("type"),
"path": fmeta.get("path"),
"minio_url": fmeta.get("minio_url"),
"markdown_file": fmeta.get("markdown_file"),
"status": fmeta.get("status"),
"content_hash": fmeta.get("content_hash"),
"file_size": fmeta.get("size") or fmeta.get("file_size"),
"content_type": fmeta.get("content_type"),
"processing_params": fmeta.get("processing_params"),
"is_folder": bool(fmeta.get("is_folder", False)),
"error_message": fmeta.get("error") or fmeta.get("error_message"),
"created_by": str(fmeta.get("created_by")) if fmeta.get("created_by") else None,
"updated_by": str(fmeta.get("updated_by")) if fmeta.get("updated_by") else None,
"created_at": _utc_dt(fmeta.get("created_at")),
"updated_at": _utc_dt(fmeta.get("updated_at")) or _utc_dt(fmeta.get("created_at")),
},
)
)
for db_id, bmap in benchmarks_meta.items():
if not isinstance(bmap, dict):
continue
for benchmark_id, bmeta in bmap.items():
benchmark_rows.append(
{
"benchmark_id": benchmark_id,
"db_id": db_id,
"name": bmeta.get("name") or benchmark_id,
"description": bmeta.get("description"),
"question_count": int(bmeta.get("question_count") or 0),
"has_gold_chunks": bool(bmeta.get("has_gold_chunks")),
"has_gold_answers": bool(bmeta.get("has_gold_answers")),
"data_file_path": bmeta.get("benchmark_file") or bmeta.get("data_file_path"),
"created_by": str(bmeta.get("created_by")) if bmeta.get("created_by") else None,
"created_at": _utc_dt(bmeta.get("created_at")),
"updated_at": _utc_dt(bmeta.get("updated_at")) or _utc_dt(bmeta.get("created_at")),
}
)
for db_id in databases_meta.keys():
result_dir = os.path.join(kb_dir, db_id, "results")
if not os.path.isdir(result_dir):
continue
for result_path in glob.glob(os.path.join(result_dir, "*.json")):
try:
data = _load_json(result_path)
except Exception as exc:
logger.warning(f"Skip invalid result file {result_path}: {exc}")
continue
task_id = data.get("task_id") or os.path.splitext(os.path.basename(result_path))[0]
benchmark_id = data.get("benchmark_id")
started_at = _utc_dt(data.get("started_at"))
result_rows.append(
{
"task_id": task_id,
"db_id": db_id,
"benchmark_id": benchmark_id,
"status": data.get("status") or "completed",
"retrieval_config": data.get("retrieval_config") or {},
"metrics": data.get("metrics") or {},
"overall_score": data.get("overall_score"),
"total_questions": int(data.get("total_questions") or 0),
"completed_questions": int(data.get("completed_questions") or 0),
"started_at": started_at,
"completed_at": _utc_dt(data.get("completed_at")) or started_at,
"created_by": str(data.get("created_by")) if data.get("created_by") else None,
}
)
interim = data.get("interim_results") or data.get("results") or []
for idx, item in enumerate(interim):
result_detail_rows.append(
(
task_id,
idx,
{
"query_text": item.get("query") or item.get("query_text") or "",
"gold_chunk_ids": item.get("gold_chunk_ids"),
"gold_answer": item.get("gold_answer"),
"generated_answer": item.get("generated_answer"),
"retrieved_chunks": item.get("retrieved_chunks"),
"metrics": item.get("metrics") or {},
},
)
)
tasks_json_path = os.path.join(config.save_dir, "tasks", "tasks.json")
task_rows: list[dict[str, Any]] = _load_json(tasks_json_path).get("tasks", []) or []
logger.info(
f"Prepared: knowledge_bases={len(kb_rows)}, knowledge_files={len(file_rows)}, "
f"benchmarks={len(benchmark_rows)}, results={len(result_rows)}, result_details={len(result_detail_rows)}, "
f"tasks={len(task_rows)}"
)
if dry_run and not execute:
return
for payload in kb_rows:
db_id = payload["db_id"]
existing = await kb_repo.get_by_id(db_id)
data = payload.copy()
if existing is None:
await kb_repo.create(data)
else:
await kb_repo.update(db_id, data)
# 先插入文件夹,再插入普通文件(确保父文件夹先存在)
folders = [(fid, data) for fid, data in file_rows if data.get("is_folder")]
files = [(fid, data) for fid, data in file_rows if not data.get("is_folder")]
for file_id, data in folders:
await file_repo.upsert(file_id=file_id, data=data)
for file_id, data in files:
await file_repo.upsert(file_id=file_id, data=data)
for payload in benchmark_rows:
# 检查知识库是否存在
kb = await kb_repo.get_by_id(payload["db_id"])
if kb is None:
logger.warning(f"Skipping benchmark {payload['benchmark_id']}: knowledge base {payload['db_id']} not found")
continue
existing = await eval_repo.get_benchmark(payload["benchmark_id"])
if existing is None:
await eval_repo.create_benchmark(payload)
for payload in result_rows:
# 检查知识库是否存在
kb = await kb_repo.get_by_id(payload["db_id"])
if kb is None:
logger.warning(f"Skipping result {payload['task_id']}: knowledge base {payload['db_id']} not found")
continue
existing = await eval_repo.get_result(payload["task_id"])
if existing is None:
await eval_repo.create_result(payload)
else:
await eval_repo.update_result(payload["task_id"], payload)
for task_id, idx, data in result_detail_rows:
await eval_repo.upsert_result_detail(task_id=task_id, query_index=idx, data=data)
for item in task_rows:
task_id = item.get("id")
if not task_id:
continue
payload = item.get("payload") or {}
result = item.get("result")
await task_repo.upsert(
task_id,
{
"name": item.get("name") or "Unnamed Task",
"type": item.get("type") or "general",
"status": item.get("status") or "pending",
"progress": float(item.get("progress") or 0.0),
"message": item.get("message") or "",
"payload": payload,
"result": result,
"error": item.get("error"),
"cancel_requested": 1 if item.get("cancel_requested") else 0,
"created_at": _utc_dt(item.get("created_at")),
"updated_at": _utc_dt(item.get("updated_at")) or _utc_dt(item.get("created_at")),
"started_at": _utc_dt(item.get("started_at")),
"completed_at": _utc_dt(item.get("completed_at")),
},
)
logger.info("Migration completed")
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--dry-run", action="store_true")
parser.add_argument("--execute", action="store_true")
parser.add_argument("--rollback", action="store_true")
args = parser.parse_args()
if not args.dry_run and not args.execute and not args.rollback:
args.dry_run = True
asyncio.run(migrate(dry_run=args.dry_run, execute=args.execute, rollback=args.rollback))
if __name__ == "__main__":
main()