style: auto-format with ruff [skip ci]

This commit is contained in:
GitHub Actions 2026-01-24 02:57:26 +00:00
parent 8443513dad
commit ec64ac795e

View File

@ -38,7 +38,8 @@ import os
import sys import sys
from dataclasses import dataclass, field from dataclasses import dataclass, field
from datetime import datetime, UTC 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__))) sys.path.insert(0, os.path.dirname(os.path.dirname(__file__)))
@ -59,15 +60,6 @@ from src.storage.postgres.models_business import (
OperationLog, OperationLog,
MessageFeedback, MessageFeedback,
MCPServer, MCPServer,
AgentConfig,
TaskRecord,
)
from src.storage.postgres.models_knowledge import (
KnowledgeBase,
KnowledgeFile,
EvaluationBenchmark,
EvaluationResult,
EvaluationResultDetail,
) )
from src.utils import logger from src.utils import logger
@ -76,9 +68,11 @@ from src.utils import logger
# 迁移阶段定义 # 迁移阶段定义
# ============================================================ # ============================================================
@dataclass @dataclass
class MigrationStage: class MigrationStage:
"""迁移阶段""" """迁移阶段"""
name: str # 阶段名称 name: str # 阶段名称
description: str # 阶段描述 description: str # 阶段描述
migrate_fn: Callable # 迁移函数 migrate_fn: Callable # 迁移函数
@ -90,6 +84,7 @@ class MigrationStage:
@dataclass @dataclass
class MigrationResult: class MigrationResult:
"""迁移结果""" """迁移结果"""
stage_name: str stage_name: str
success: bool success: bool
dry_run: bool dry_run: bool
@ -240,6 +235,7 @@ class SqliteMCPServer(SQLiteBase):
# 工具函数 # 工具函数
# ============================================================ # ============================================================
def _utc_dt(value: Any) -> datetime | None: def _utc_dt(value: Any) -> datetime | None:
"""转换各种 datetime 格式为 naive UTC datetime""" """转换各种 datetime 格式为 naive UTC datetime"""
if not value: if not value:
@ -297,6 +293,7 @@ def _log_separator(title: str = "", char: str = "=", width: int = 60) -> str:
# SQLite 读取器 # SQLite 读取器
# ============================================================ # ============================================================
class SQLiteReader: class SQLiteReader:
"""SQLite 数据读取器""" """SQLite 数据读取器"""
@ -324,6 +321,7 @@ class SQLiteReader:
# 迁移阶段实现 # 迁移阶段实现
# ============================================================ # ============================================================
class MigrationRunner: class MigrationRunner:
"""迁移运行器""" """迁移运行器"""
@ -687,7 +685,8 @@ class MigrationRunner:
kb_rows = [] kb_rows = []
kb_type_dirs = [ 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" if os.path.isdir(p) and os.path.basename(p) != "uploads"
] ]
@ -701,21 +700,23 @@ class MigrationRunner:
g = global_meta.get(db_id, {}) g = global_meta.get(db_id, {})
created_at = _utc_dt(g.get("created_at") or db_meta.get("created_at")) 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 updated_at = _utc_dt(g.get("updated_at")) or created_at
kb_rows.append({ kb_rows.append(
"db_id": db_id, {
"name": g.get("name") or db_meta.get("name") or db_id, "db_id": db_id,
"description": g.get("description") or db_meta.get("description"), "name": g.get("name") or db_meta.get("name") or db_id,
"kb_type": g.get("kb_type") or db_meta.get("kb_type") or kb_type, "description": g.get("description") or db_meta.get("description"),
"embed_info": db_meta.get("embed_info") or g.get("embed_info"), "kb_type": g.get("kb_type") or db_meta.get("kb_type") or kb_type,
"llm_info": db_meta.get("llm_info") or g.get("llm_info"), "embed_info": db_meta.get("embed_info") or g.get("embed_info"),
"query_params": db_meta.get("query_params") or g.get("query_params"), "llm_info": db_meta.get("llm_info") or g.get("llm_info"),
"additional_params": g.get("additional_params") or db_meta.get("metadata") or {}, "query_params": db_meta.get("query_params") or g.get("query_params"),
"share_config": {"is_shared": True, "accessible_departments": []}, "additional_params": g.get("additional_params") or db_meta.get("metadata") or {},
"mindmap": g.get("mindmap"), "share_config": {"is_shared": True, "accessible_departments": []},
"sample_questions": g.get("sample_questions") or [], "mindmap": g.get("mindmap"),
"created_at": created_at, "sample_questions": g.get("sample_questions") or [],
"updated_at": updated_at, "created_at": created_at,
}) "updated_at": updated_at,
}
)
result.records_total = len(kb_rows) result.records_total = len(kb_rows)
@ -725,6 +726,7 @@ class MigrationRunner:
return result return result
from src.repositories.knowledge_base_repository import KnowledgeBaseRepository from src.repositories.knowledge_base_repository import KnowledgeBaseRepository
kb_repo = KnowledgeBaseRepository() kb_repo = KnowledgeBaseRepository()
for payload in kb_rows: for payload in kb_rows:
@ -744,7 +746,8 @@ class MigrationRunner:
file_rows = [] file_rows = []
kb_type_dirs = [ 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" if os.path.isdir(p) and os.path.basename(p) != "uploads"
] ]
@ -757,28 +760,30 @@ class MigrationRunner:
db_id = fmeta.get("database_id") db_id = fmeta.get("database_id")
if not db_id: if not db_id:
continue continue
file_rows.append({ file_rows.append(
"file_id": file_id, {
"db_id": db_id, "file_id": file_id,
"parent_id": fmeta.get("parent_id"), "db_id": db_id,
"filename": fmeta.get("filename") or "", "parent_id": fmeta.get("parent_id"),
"original_filename": fmeta.get("original_filename") or fmeta.get("file_name"), "filename": fmeta.get("filename") or "",
"file_type": fmeta.get("file_type") or fmeta.get("type"), "original_filename": fmeta.get("original_filename") or fmeta.get("file_name"),
"path": fmeta.get("path"), "file_type": fmeta.get("file_type") or fmeta.get("type"),
"minio_url": fmeta.get("minio_url"), "path": fmeta.get("path"),
"markdown_file": fmeta.get("markdown_file"), "minio_url": fmeta.get("minio_url"),
"status": fmeta.get("status"), "markdown_file": fmeta.get("markdown_file"),
"content_hash": fmeta.get("content_hash"), "status": fmeta.get("status"),
"file_size": fmeta.get("size") or fmeta.get("file_size"), "content_hash": fmeta.get("content_hash"),
"content_type": fmeta.get("content_type"), "file_size": fmeta.get("size") or fmeta.get("file_size"),
"processing_params": fmeta.get("processing_params"), "content_type": fmeta.get("content_type"),
"is_folder": bool(fmeta.get("is_folder", False)), "processing_params": fmeta.get("processing_params"),
"error_message": fmeta.get("error") or fmeta.get("error_message"), "is_folder": bool(fmeta.get("is_folder", False)),
"created_by": str(fmeta.get("created_by")) if fmeta.get("created_by") else None, "error_message": fmeta.get("error") or fmeta.get("error_message"),
"updated_by": str(fmeta.get("updated_by")) if fmeta.get("updated_by") else None, "created_by": str(fmeta.get("created_by")) if fmeta.get("created_by") else None,
"created_at": _utc_dt(fmeta.get("created_at")), "updated_by": str(fmeta.get("updated_by")) if fmeta.get("updated_by") else None,
"updated_at": _utc_dt(fmeta.get("updated_at")) or _utc_dt(fmeta.get("created_at")), "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) result.records_total = len(file_rows)
@ -789,6 +794,7 @@ class MigrationRunner:
return result return result
from src.repositories.knowledge_file_repository import KnowledgeFileRepository from src.repositories.knowledge_file_repository import KnowledgeFileRepository
file_repo = KnowledgeFileRepository() file_repo = KnowledgeFileRepository()
# 先插入文件夹 # 先插入文件夹
@ -815,12 +821,14 @@ class MigrationRunner:
total_migrated = 0 total_migrated = 0
kb_type_dirs = [ 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" if os.path.isdir(p) and os.path.basename(p) != "uploads"
] ]
from src.repositories.evaluation_repository import EvaluationRepository from src.repositories.evaluation_repository import EvaluationRepository
from src.repositories.knowledge_base_repository import KnowledgeBaseRepository from src.repositories.knowledge_base_repository import KnowledgeBaseRepository
eval_repo = EvaluationRepository() eval_repo = EvaluationRepository()
kb_repo = KnowledgeBaseRepository() kb_repo = KnowledgeBaseRepository()
@ -836,19 +844,21 @@ class MigrationRunner:
if not isinstance(bmap, dict): if not isinstance(bmap, dict):
continue continue
for benchmark_id, bmeta in bmap.items(): for benchmark_id, bmeta in bmap.items():
benchmark_rows.append({ benchmark_rows.append(
"benchmark_id": benchmark_id, {
"db_id": db_id, "benchmark_id": benchmark_id,
"name": bmeta.get("name") or benchmark_id, "db_id": db_id,
"description": bmeta.get("description"), "name": bmeta.get("name") or benchmark_id,
"question_count": int(bmeta.get("question_count") or 0), "description": bmeta.get("description"),
"has_gold_chunks": bool(bmeta.get("has_gold_chunks")), "question_count": int(bmeta.get("question_count") or 0),
"has_gold_answers": bool(bmeta.get("has_gold_answers")), "has_gold_chunks": bool(bmeta.get("has_gold_chunks")),
"data_file_path": bmeta.get("benchmark_file") or bmeta.get("data_file_path"), "has_gold_answers": bool(bmeta.get("has_gold_answers")),
"created_by": str(bmeta.get("created_by")) if bmeta.get("created_by") else None, "data_file_path": bmeta.get("benchmark_file") or bmeta.get("data_file_path"),
"created_at": _utc_dt(bmeta.get("created_at")), "created_by": str(bmeta.get("created_by")) if bmeta.get("created_by") else None,
"updated_at": _utc_dt(bmeta.get("updated_at")) or _utc_dt(bmeta.get("created_at")), "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) 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] task_id = data.get("task_id") or os.path.splitext(os.path.basename(result_path))[0]
benchmark_id = data.get("benchmark_id") benchmark_id = data.get("benchmark_id")
started_at = _utc_dt(data.get("started_at")) started_at = _utc_dt(data.get("started_at"))
result_rows.append({ result_rows.append(
"task_id": task_id, {
"db_id": db_id, "task_id": task_id,
"benchmark_id": benchmark_id, "db_id": db_id,
"status": data.get("status") or "completed", "benchmark_id": benchmark_id,
"retrieval_config": data.get("retrieval_config") or {}, "status": data.get("status") or "completed",
"metrics": data.get("metrics") or {}, "retrieval_config": data.get("retrieval_config") or {},
"overall_score": data.get("overall_score"), "metrics": data.get("metrics") or {},
"total_questions": int(data.get("total_questions") or 0), "overall_score": data.get("overall_score"),
"completed_questions": int(data.get("completed_questions") or 0), "total_questions": int(data.get("total_questions") or 0),
"started_at": started_at, "completed_questions": int(data.get("completed_questions") or 0),
"completed_at": _utc_dt(data.get("completed_at")) or started_at, "started_at": started_at,
"created_by": str(data.get("created_by")) if data.get("created_by") else None, "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 [] interim = data.get("interim_results") or data.get("results") or []
for idx, item in enumerate(interim): for idx, item in enumerate(interim):
result_detail_rows.append({ result_detail_rows.append(
"task_id": task_id, {
"query_index": idx, "task_id": task_id,
"query_text": item.get("query") or item.get("query_text") or "", "query_index": idx,
"gold_chunk_ids": item.get("gold_chunk_ids"), "query_text": item.get("query") or item.get("query_text") or "",
"gold_answer": item.get("gold_answer"), "gold_chunk_ids": item.get("gold_chunk_ids"),
"generated_answer": item.get("generated_answer"), "gold_answer": item.get("gold_answer"),
"retrieved_chunks": item.get("retrieved_chunks"), "generated_answer": item.get("generated_answer"),
"metrics": item.get("metrics") or {}, "retrieved_chunks": item.get("retrieved_chunks"),
}) "metrics": item.get("metrics") or {},
}
)
result.records_total += len(result_rows) + len(result_detail_rows) result.records_total += len(result_rows) + len(result_detail_rows)
@ -968,6 +982,7 @@ class MigrationRunner:
return result return result
from src.repositories.task_repository import TaskRepository from src.repositories.task_repository import TaskRepository
task_repo = TaskRepository() task_repo = TaskRepository()
for item in task_items: for item in task_items:
@ -1147,6 +1162,7 @@ class MigrationRunner:
async with pg_manager.get_async_session_context() as session: async with pg_manager.get_async_session_context() as session:
from sqlalchemy import func from sqlalchemy import func
result = await session.execute(select(func.count(getattr(pg_model, pk_column)))) result = await session.execute(select(func.count(getattr(pg_model, pk_column))))
pg_count = result.scalar() or 0 pg_count = result.scalar() or 0
@ -1169,7 +1185,8 @@ class MigrationRunner:
json_file_count = 0 json_file_count = 0
kb_type_dirs = [ 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" if os.path.isdir(p) and os.path.basename(p) != "uploads"
] ]
@ -1197,7 +1214,11 @@ class MigrationRunner:
pg_file_count = len(all_files) 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_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_kb = "" if results["knowledge_bases"]["match"] else ""
status_file = "" if results["knowledge_files"]["match"] else "" status_file = "" if results["knowledge_files"]["match"] else ""
@ -1212,6 +1233,7 @@ class MigrationRunner:
# 阶段定义 # 阶段定义
# ============================================================ # ============================================================
def get_stages() -> dict[str, MigrationStage]: def get_stages() -> dict[str, MigrationStage]:
"""获取所有迁移阶段""" """获取所有迁移阶段"""
runner = MigrationRunner() runner = MigrationRunner()
@ -1328,6 +1350,7 @@ def get_stage_groups() -> dict[str, list[str]]:
# 主函数 # 主函数
# ============================================================ # ============================================================
async def main() -> None: async def main() -> None:
parser = argparse.ArgumentParser(description="统一数据迁移脚本") parser = argparse.ArgumentParser(description="统一数据迁移脚本")
parser.add_argument("--dry-run", action="store_true", help="预览迁移,不执行") parser.add_argument("--dry-run", action="store_true", help="预览迁移,不执行")