From ec64ac795edfeb3fd2b4b09ac71bda9ae104e59b Mon Sep 17 00:00:00 2001 From: GitHub Actions Date: Sat, 24 Jan 2026 02:57:26 +0000 Subject: [PATCH] style: auto-format with ruff [skip ci] --- scripts/migrate_all.py | 201 +++++++++++++++++++++++------------------ 1 file changed, 112 insertions(+), 89 deletions(-) diff --git a/scripts/migrate_all.py b/scripts/migrate_all.py index c80e7cf9..1dca7acb 100644 --- a/scripts/migrate_all.py +++ b/scripts/migrate_all.py @@ -38,7 +38,8 @@ import os import sys from dataclasses import dataclass, field from datetime import datetime, UTC -from typing import Any, Callable +from typing import Any +from collections.abc import Callable # 确保路径正确 sys.path.insert(0, os.path.dirname(os.path.dirname(__file__))) @@ -59,15 +60,6 @@ from src.storage.postgres.models_business import ( OperationLog, MessageFeedback, MCPServer, - AgentConfig, - TaskRecord, -) -from src.storage.postgres.models_knowledge import ( - KnowledgeBase, - KnowledgeFile, - EvaluationBenchmark, - EvaluationResult, - EvaluationResultDetail, ) from src.utils import logger @@ -76,9 +68,11 @@ from src.utils import logger # 迁移阶段定义 # ============================================================ + @dataclass class MigrationStage: """迁移阶段""" + name: str # 阶段名称 description: str # 阶段描述 migrate_fn: Callable # 迁移函数 @@ -90,6 +84,7 @@ class MigrationStage: @dataclass class MigrationResult: """迁移结果""" + stage_name: str success: bool dry_run: bool @@ -240,6 +235,7 @@ class SqliteMCPServer(SQLiteBase): # 工具函数 # ============================================================ + def _utc_dt(value: Any) -> datetime | None: """转换各种 datetime 格式为 naive UTC datetime""" if not value: @@ -297,6 +293,7 @@ def _log_separator(title: str = "", char: str = "=", width: int = 60) -> str: # SQLite 读取器 # ============================================================ + class SQLiteReader: """SQLite 数据读取器""" @@ -324,6 +321,7 @@ class SQLiteReader: # 迁移阶段实现 # ============================================================ + class MigrationRunner: """迁移运行器""" @@ -687,7 +685,8 @@ class MigrationRunner: kb_rows = [] kb_type_dirs = [ - p for p in glob.glob(os.path.join(base_dir, "*_data")) + p + for p in glob.glob(os.path.join(base_dir, "*_data")) if os.path.isdir(p) and os.path.basename(p) != "uploads" ] @@ -701,21 +700,23 @@ class MigrationRunner: 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": {"is_shared": True, "accessible_departments": []}, - "mindmap": g.get("mindmap"), - "sample_questions": g.get("sample_questions") or [], - "created_at": created_at, - "updated_at": updated_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": {"is_shared": True, "accessible_departments": []}, + "mindmap": g.get("mindmap"), + "sample_questions": g.get("sample_questions") or [], + "created_at": created_at, + "updated_at": updated_at, + } + ) result.records_total = len(kb_rows) @@ -725,6 +726,7 @@ class MigrationRunner: return result from src.repositories.knowledge_base_repository import KnowledgeBaseRepository + kb_repo = KnowledgeBaseRepository() for payload in kb_rows: @@ -744,7 +746,8 @@ class MigrationRunner: file_rows = [] kb_type_dirs = [ - p for p in glob.glob(os.path.join(base_dir, "*_data")) + p + for p in glob.glob(os.path.join(base_dir, "*_data")) if os.path.isdir(p) and os.path.basename(p) != "uploads" ] @@ -757,28 +760,30 @@ class MigrationRunner: db_id = fmeta.get("database_id") if not db_id: continue - file_rows.append({ - "file_id": 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")), - }) + file_rows.append( + { + "file_id": 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")), + } + ) result.records_total = len(file_rows) @@ -789,6 +794,7 @@ class MigrationRunner: return result from src.repositories.knowledge_file_repository import KnowledgeFileRepository + file_repo = KnowledgeFileRepository() # 先插入文件夹 @@ -815,12 +821,14 @@ class MigrationRunner: total_migrated = 0 kb_type_dirs = [ - p for p in glob.glob(os.path.join(base_dir, "*_data")) + p + for p in glob.glob(os.path.join(base_dir, "*_data")) if os.path.isdir(p) and os.path.basename(p) != "uploads" ] from src.repositories.evaluation_repository import EvaluationRepository from src.repositories.knowledge_base_repository import KnowledgeBaseRepository + eval_repo = EvaluationRepository() kb_repo = KnowledgeBaseRepository() @@ -836,19 +844,21 @@ class MigrationRunner: 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")), - }) + 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")), + } + ) result.records_total += len(benchmark_rows) @@ -892,32 +902,36 @@ class MigrationRunner: 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, - }) + 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": task_id, - "query_index": 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 {}, - }) + result_detail_rows.append( + { + "task_id": task_id, + "query_index": 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 {}, + } + ) result.records_total += len(result_rows) + len(result_detail_rows) @@ -968,6 +982,7 @@ class MigrationRunner: return result from src.repositories.task_repository import TaskRepository + task_repo = TaskRepository() for item in task_items: @@ -1147,6 +1162,7 @@ class MigrationRunner: async with pg_manager.get_async_session_context() as session: from sqlalchemy import func + result = await session.execute(select(func.count(getattr(pg_model, pk_column)))) pg_count = result.scalar() or 0 @@ -1169,7 +1185,8 @@ class MigrationRunner: json_file_count = 0 kb_type_dirs = [ - p for p in glob.glob(os.path.join(base_dir, "*_data")) + p + for p in glob.glob(os.path.join(base_dir, "*_data")) if os.path.isdir(p) and os.path.basename(p) != "uploads" ] @@ -1197,7 +1214,11 @@ class MigrationRunner: pg_file_count = len(all_files) results["knowledge_bases"] = {"json": json_kb_count, "pg": pg_kb_count, "match": json_kb_count == pg_kb_count} - results["knowledge_files"] = {"json": json_file_count, "pg": pg_file_count, "match": json_file_count == pg_file_count} + results["knowledge_files"] = { + "json": json_file_count, + "pg": pg_file_count, + "match": json_file_count == pg_file_count, + } status_kb = "✅" if results["knowledge_bases"]["match"] else "❌" status_file = "✅" if results["knowledge_files"]["match"] else "❌" @@ -1212,6 +1233,7 @@ class MigrationRunner: # 阶段定义 # ============================================================ + def get_stages() -> dict[str, MigrationStage]: """获取所有迁移阶段""" runner = MigrationRunner() @@ -1328,6 +1350,7 @@ def get_stage_groups() -> dict[str, list[str]]: # 主函数 # ============================================================ + async def main() -> None: parser = argparse.ArgumentParser(description="统一数据迁移脚本") parser.add_argument("--dry-run", action="store_true", help="预览迁移,不执行")