From 3f96c9c6502243b06efd4bb34a189e72e3a5557c Mon Sep 17 00:00:00 2001
From: yearning <10538594+wangweifeng1999@user.noreply.gitee.com>
Date: 星期一, 28 九月 2026 19:24:55 +0800
Subject: [PATCH] 复制项目
---
backend/app/import_service.py | 501 +++++++++++++++++++++++++++++++++++++++++++++++++++++++
1 files changed, 501 insertions(+), 0 deletions(-)
diff --git a/backend/app/import_service.py b/backend/app/import_service.py
new file mode 100644
index 0000000..0e8f72b
--- /dev/null
+++ b/backend/app/import_service.py
@@ -0,0 +1,501 @@
+import hashlib
+import io
+import json
+import zipfile
+from pathlib import Path
+from typing import Any
+
+from fastapi import HTTPException
+from sqlalchemy import delete, func, select
+from sqlalchemy.orm import Session
+
+from app.grading import answers_json_to_map, grade_item
+from app.models import (
+ ExamPaperMap,
+ ExaminationExamExaminee,
+ ExaminationExaminee,
+ ExaminationExamineeExamAnswer,
+ ExaminationExamineeExamAnswerItem,
+ ExaminationPaperItem,
+ ExaminationPaperSource,
+ ImportBatch,
+)
+from app.sql_parser import find_file_content, rows_from_sql
+
+
+def load_zip_files(zip_bytes: bytes) -> dict[str, bytes]:
+ out: dict[str, bytes] = {}
+ with zipfile.ZipFile(io.BytesIO(zip_bytes)) as zf:
+ for info in zf.infolist():
+ if info.is_dir():
+ continue
+ if info.filename.lower().endswith(".sql"):
+ out[info.filename] = zf.read(info)
+ return out
+
+
+def load_dir_sql(dir_path: Path) -> dict[str, bytes]:
+ out: dict[str, bytes] = {}
+ for p in dir_path.rglob("*.sql"):
+ out[str(p.relative_to(dir_path))] = p.read_bytes()
+ return out
+
+
+def _require_sql_kind(files: dict[str, bytes], kind: str, label: str) -> str:
+ text = find_file_content(files, kind)
+ if not text:
+ raise HTTPException(400, detail=f"ZIP 缂哄皯 {label} SQL 鏂囦欢")
+ return text
+
+
+def _paper_identity(row: dict[str, Any]) -> tuple[str, str]:
+ return (str(row.get("external_paper_id", "")), str(row.get("revision", "")))
+
+
+def _collect_exam_ids(exam_rows: list[dict[str, Any]]) -> list[int]:
+ ids = sorted({int(r["exam_id"]) for r in exam_rows if r.get("exam_id") is not None})
+ if not ids:
+ raise HTTPException(400, detail="鑰冭瘯鑰冪敓 SQL 涓湭鎵惧埌鏈夋晥鐨� exam_id")
+ return ids
+
+
+def _parse_exam_to_paper_sql_map(exam_sql: str) -> dict[int, int]:
+ out: dict[int, int] = {}
+ for r in rows_from_sql(exam_sql):
+ eid = r.get("id")
+ if eid is None:
+ continue
+ psid = r.get("paper_source_id")
+ if psid is None:
+ psid = r.get("paper_id")
+ if psid is None:
+ continue
+ out[int(eid)] = int(psid)
+ return out
+
+
+def _resolve_exam_paper_map(
+ exam_ids: list[int],
+ paper_rows: list[dict[str, Any]],
+ files: dict[str, bytes],
+) -> dict[int, int]:
+ """exam_id -> paper_source_id锛圫QL 鏂囦欢涓殑璇曞嵎 id锛夈��"""
+ keys = {_paper_identity(r) for r in paper_rows}
+ sql_paper_ids = {int(r["id"]) for r in paper_rows if r.get("id") is not None}
+
+ if len(keys) <= 1:
+ if not paper_rows:
+ raise HTTPException(400, detail="璇曞嵎 SQL 鏃犳湁鏁堟暟鎹�")
+ canonical = int(paper_rows[0]["id"])
+ return {eid: canonical for eid in exam_ids}
+
+ exam_text = find_file_content(files, "exam")
+ if not exam_text:
+ raise HTTPException(
+ 400,
+ detail="ZIP 鍚浠借瘯鍗锋椂锛屽繀椤绘彁渚� examination_exam SQL锛岀敤浜庡叧鑱旀瘡鍦鸿�冭瘯鐨勮瘯鍗�",
+ )
+ meta = _parse_exam_to_paper_sql_map(exam_text)
+ if not meta:
+ raise HTTPException(400, detail="examination_exam SQL 鏃犳湁鏁堟暟鎹垨鏈寘鍚� paper_source_id")
+
+ result: dict[int, int] = {}
+ for eid in exam_ids:
+ ps_sql = meta.get(eid)
+ if ps_sql is None:
+ raise HTTPException(400, detail=f"examination_exam 涓己灏� exam_id={eid} 鐨勮瘯鍗峰叧鑱�")
+ if ps_sql not in sql_paper_ids:
+ raise HTTPException(
+ 400,
+ detail=f"exam_id={eid} 鍏宠仈鐨� paper_source_id={ps_sql} 涓嶅湪 ZIP 璇曞嵎 SQL 涓�",
+ )
+ result[eid] = ps_sql
+ return result
+
+
+def _get_or_create_paper_source(
+ db: Session,
+ batch: ImportBatch,
+ row: dict[str, Any],
+) -> int:
+ ext_id, revision = _paper_identity(row)
+ existing = db.scalar(
+ select(ExaminationPaperSource.id).where(
+ ExaminationPaperSource.external_paper_id == ext_id,
+ ExaminationPaperSource.revision == revision,
+ )
+ )
+ if existing:
+ return int(existing)
+
+ ps = ExaminationPaperSource(
+ id=int(row["id"]),
+ external_paper_id=ext_id,
+ revision=revision,
+ sha256=str(row.get("sha256") or ""),
+ paper_name=str(row.get("paper_name") or ""),
+ occupation=str(row.get("occupation") or ""),
+ level=str(row.get("level") or ""),
+ subject=int(row["subject"]) if row.get("subject") is not None else None,
+ struct_json=str(row.get("struct_json") or ""),
+ item_count=int(row["item_count"]) if row.get("item_count") is not None else None,
+ import_batch_id=batch.id,
+ )
+ db.add(ps)
+ db.flush()
+ return int(ps.id)
+
+
+def _items_for_grading(
+ db: Session,
+ batch: ImportBatch,
+ sql_ps_id: int,
+ db_ps_id: int,
+ item_rows: list[dict[str, Any]],
+) -> dict[str, ExaminationPaperItem]:
+ ps = db.get(ExaminationPaperSource, db_ps_id)
+ if not ps:
+ raise HTTPException(500, detail=f"璇曞嵎 {db_ps_id} 涓嶅瓨鍦�")
+
+ if ps.import_batch_id != batch.id:
+ items = list(
+ db.scalars(
+ select(ExaminationPaperItem).where(ExaminationPaperItem.paper_source_id == db_ps_id)
+ )
+ )
+ return {it.external_item_id: it for it in items}
+
+ items_by_ext: dict[str, ExaminationPaperItem] = {}
+ total_max = 0.0
+ for r in item_rows:
+ if int(r.get("paper_source_id") or 0) != sql_ps_id:
+ continue
+ ms = float(r.get("max_score") or 1)
+ total_max += ms
+ it = ExaminationPaperItem(
+ id=int(r["id"]),
+ paper_source_id=db_ps_id,
+ external_item_id=str(r["external_item_id"]),
+ section_ref=str(r.get("section_ref") or ""),
+ part_ref=str(r.get("part_ref") or ""),
+ part_name=str(r.get("part_name") or ""),
+ item_type=int(r["item_type"]) if r.get("item_type") is not None else None,
+ item_order=int(r["item_order"]) if r.get("item_order") is not None else None,
+ title_html=str(r.get("title_html") or ""),
+ detail_json=str(r.get("detail_json") or ""),
+ max_score=ms,
+ reference_answer=str(r.get("reference_answer") or ""),
+ answer_rules=str(r.get("answer_rules") or ""),
+ )
+ db.add(it)
+ items_by_ext[it.external_item_id] = it
+ ps.total_score = total_max
+ return items_by_ext
+
+
+def run_import(
+ db: Session,
+ subject_name: str,
+ zip_filename: str,
+ file_bytes: bytes,
+ user_id: int | None,
+) -> dict[str, Any]:
+ files = load_zip_files(file_bytes)
+ sha = hashlib.sha256(file_bytes).hexdigest()
+
+ examinee_sql = _require_sql_kind(files, "examinee", "鑰冪敓")
+ exam_examinee_sql = _require_sql_kind(files, "exam_examinee", "鑰冭瘯鑰冪敓")
+ paper_sql = _require_sql_kind(files, "paper_source", "璇曞嵎")
+ item_sql = _require_sql_kind(files, "paper_item", "璇曢")
+ answer_sql = _require_sql_kind(files, "answer", "鑰冪敓绛旀")
+
+ paper_rows = rows_from_sql(paper_sql)
+ if not paper_rows:
+ raise HTTPException(400, detail="璇曞嵎 SQL 鏃犳湁鏁堟暟鎹�")
+
+ exam_rows = rows_from_sql(exam_examinee_sql)
+ if not exam_rows:
+ raise HTTPException(400, detail="鑰冭瘯鑰冪敓 SQL 鏃犳湁鏁堟暟鎹�")
+
+ exam_ids = _collect_exam_ids(exam_rows)
+ for eid in exam_ids:
+ if db.scalar(select(ExamPaperMap.id).where(ExamPaperMap.exam_id == eid)):
+ raise HTTPException(409, detail=f"exam_id {eid} 宸插瓨鍦紝璇峰厛娓呴櫎瀵瑰簲瀵煎叆鎵规")
+
+ exam_to_sql_paper = _resolve_exam_paper_map(exam_ids, paper_rows, files)
+
+ paper_row_by_sql_id = {int(r["id"]): r for r in paper_rows if r.get("id") is not None}
+ needed_sql_paper_ids = {exam_to_sql_paper[e] for e in exam_ids}
+ for sql_id in needed_sql_paper_ids:
+ if sql_id not in paper_row_by_sql_id:
+ raise HTTPException(400, detail=f"璇曞嵎 SQL 缂哄皯 id={sql_id}")
+
+ warnings: list[dict[str, Any]] = []
+ report: dict[str, Any] = {"warnings": warnings, "counts": {}}
+
+ primary_paper_row = paper_row_by_sql_id[exam_to_sql_paper[exam_ids[0]]]
+ ext_id, revision = _paper_identity(primary_paper_row)
+
+ try:
+ batch = ImportBatch(
+ subject_name=subject_name.strip(),
+ zip_filename=zip_filename,
+ sha256=sha,
+ status="success",
+ imported_by=user_id,
+ external_paper_id=ext_id,
+ revision=revision,
+ exam_id=exam_ids[0],
+ exam_ids_json=json.dumps(exam_ids, ensure_ascii=False),
+ occupation=str(primary_paper_row.get("occupation") or ""),
+ level=str(primary_paper_row.get("level") or ""),
+ subject_code=int(primary_paper_row["subject"])
+ if primary_paper_row.get("subject") is not None
+ else None,
+ )
+ db.add(batch)
+ db.flush()
+
+ sql_to_db_paper: dict[int, int] = {}
+ for sql_id in needed_sql_paper_ids:
+ sql_to_db_paper[sql_id] = _get_or_create_paper_source(
+ db, batch, paper_row_by_sql_id[sql_id]
+ )
+
+ batch.paper_source_id = sql_to_db_paper[exam_to_sql_paper[exam_ids[0]]]
+
+ item_rows = rows_from_sql(item_sql)
+ grading_items: dict[int, dict[str, ExaminationPaperItem]] = {}
+ for sql_id in set(exam_to_sql_paper.values()):
+ db_ps = sql_to_db_paper[sql_id]
+ grading_items[db_ps] = _items_for_grading(db, batch, sql_id, db_ps, item_rows)
+
+ for r in rows_from_sql(examinee_sql):
+ eid = int(r["id"])
+ existing = db.get(ExaminationExaminee, eid)
+ if existing:
+ existing.username = r.get("username")
+ existing.nickname = r.get("nickname")
+ existing.id_card = r.get("id_card")
+ existing.mobile = r.get("mobile")
+ existing.sex = int(r["sex"]) if r.get("sex") is not None else None
+ else:
+ db.add(
+ ExaminationExaminee(
+ id=eid,
+ username=r.get("username"),
+ nickname=r.get("nickname"),
+ id_card=r.get("id_card"),
+ mobile=r.get("mobile"),
+ sex=int(r["sex"]) if r.get("sex") is not None else None,
+ deleted=int(r.get("deleted") or 0),
+ )
+ )
+ db.flush()
+
+ answer_rows = {int(r["exam_examinee_id"]): r for r in rows_from_sql(answer_sql)}
+
+ for r in exam_rows:
+ aid = int(r["id"])
+ row_exam_id = int(r["exam_id"])
+ imported_score = float(r.get("score") or 0)
+ sql_ps = exam_to_sql_paper[row_exam_id]
+ db_ps = sql_to_db_paper[sql_ps]
+ items_by_ext = grading_items[db_ps]
+
+ attempt = ExaminationExamExaminee(
+ id=aid,
+ exam_id=row_exam_id,
+ examinee_id=int(r["examinee_id"]),
+ room_id=r.get("room_id"),
+ seat_number=r.get("seat_number"),
+ status=int(r["status"]) if r.get("status") is not None else None,
+ score=0,
+ imported_score=imported_score,
+ level_str=r.get("level_str"),
+ start_time=r.get("start_time"),
+ end_time=r.get("end_time"),
+ subject=int(r["subject"]) if r.get("subject") is not None else None,
+ import_batch_id=batch.id,
+ )
+ db.add(attempt)
+
+ ar = answer_rows.get(aid)
+ if not ar:
+ raise HTTPException(400, detail=f"缂哄皯鑰冪敓绛旀 exam_examinee_id={aid}")
+ answers_str = str(ar.get("answers") or "[]")
+ db.add(
+ ExaminationExamineeExamAnswer(
+ id=int(ar["id"]),
+ exam_examinee_id=aid,
+ answers=answers_str,
+ deleted=int(ar.get("deleted") or 0),
+ )
+ )
+ amap = answers_json_to_map(answers_str)
+ calc = 0.0
+ for ext, item in items_by_ext.items():
+ ua = amap.get(ext, "")
+ ok, sc = grade_item(
+ ua,
+ item.reference_answer or "",
+ item.answer_rules,
+ item.item_type,
+ item.max_score,
+ )
+ calc += sc
+ db.add(
+ ExaminationExamineeExamAnswerItem(
+ exam_examinee_id=aid,
+ external_item_id=ext,
+ answer_raw=ua,
+ is_correct=ok,
+ score=sc,
+ )
+ )
+ attempt.score = calc
+ if abs(calc - imported_score) > 0.01:
+ warnings.append(
+ {
+ "exam_examinee_id": aid,
+ "imported_score": imported_score,
+ "calculated_score": calc,
+ }
+ )
+
+ for eid in exam_ids:
+ db.add(
+ ExamPaperMap(
+ exam_id=eid,
+ paper_source_id=sql_to_db_paper[exam_to_sql_paper[eid]],
+ import_batch_id=batch.id,
+ )
+ )
+
+ report["counts"] = {
+ "examinees": len(rows_from_sql(examinee_sql)),
+ "attempts": len(exam_rows),
+ "items": len(item_rows),
+ "exam_id": exam_ids[0],
+ "exam_ids": exam_ids,
+ "papers": len(set(sql_to_db_paper.values())),
+ }
+ batch.report_json = json.dumps(report, ensure_ascii=False)
+ db.commit()
+ db.refresh(batch)
+ return {"batch_id": batch.id, "report": report}
+ except HTTPException:
+ db.rollback()
+ raise
+ except Exception as e:
+ db.rollback()
+ raise HTTPException(500, detail=f"瀵煎叆澶辫触: {e}") from e
+
+
+def _paper_owned_by_batch(db: Session, ps_id: int, batch_id: int) -> bool:
+ ps = db.get(ExaminationPaperSource, ps_id)
+ return ps is not None and ps.import_batch_id == batch_id
+
+
+def _paper_still_referenced(db: Session, ps_id: int, batch_id: int) -> bool:
+ other_maps = db.scalar(
+ select(func.count())
+ .select_from(ExamPaperMap)
+ .where(
+ ExamPaperMap.paper_source_id == ps_id,
+ ExamPaperMap.import_batch_id != batch_id,
+ )
+ )
+ return (other_maps or 0) > 0
+
+
+def purge_batch(db: Session, batch_id: int) -> dict[str, int]:
+ batch = db.get(ImportBatch, batch_id)
+ if not batch:
+ raise HTTPException(404, detail="瀵煎叆鎵规涓嶅瓨鍦�")
+
+ counts: dict[str, int] = {}
+ attempt_ids = list(
+ db.scalars(
+ select(ExaminationExamExaminee.id).where(
+ ExaminationExamExaminee.import_batch_id == batch_id
+ )
+ )
+ )
+ examinee_ids_to_check: set[int] = set()
+ if attempt_ids:
+ examinee_ids_to_check = set(
+ db.scalars(
+ select(ExaminationExamExaminee.examinee_id).where(
+ ExaminationExamExaminee.id.in_(attempt_ids)
+ )
+ )
+ )
+ counts["answer_item"] = db.execute(
+ delete(ExaminationExamineeExamAnswerItem).where(
+ ExaminationExamineeExamAnswerItem.exam_examinee_id.in_(attempt_ids)
+ )
+ ).rowcount
+ counts["exam_answer"] = db.execute(
+ delete(ExaminationExamineeExamAnswer).where(
+ ExaminationExamineeExamAnswer.exam_examinee_id.in_(attempt_ids)
+ )
+ ).rowcount
+ counts["exam_examinee"] = db.execute(
+ delete(ExaminationExamExaminee).where(ExaminationExamExaminee.import_batch_id == batch_id)
+ ).rowcount
+
+ paper_ids_to_check = set(
+ db.scalars(
+ select(ExamPaperMap.paper_source_id).where(ExamPaperMap.import_batch_id == batch_id)
+ )
+ )
+ if batch.paper_source_id:
+ paper_ids_to_check.add(batch.paper_source_id)
+
+ counts["exam_paper_map"] = db.execute(
+ delete(ExamPaperMap).where(ExamPaperMap.import_batch_id == batch_id)
+ ).rowcount
+
+ for ps_id in paper_ids_to_check:
+ if not _paper_owned_by_batch(db, ps_id, batch_id):
+ continue
+ if _paper_still_referenced(db, ps_id, batch_id):
+ continue
+ counts["paper_item"] = counts.get("paper_item", 0) + (
+ db.execute(
+ delete(ExaminationPaperItem).where(ExaminationPaperItem.paper_source_id == ps_id)
+ ).rowcount
+ or 0
+ )
+ counts["paper_source"] = counts.get("paper_source", 0) + (
+ db.execute(delete(ExaminationPaperSource).where(ExaminationPaperSource.id == ps_id)).rowcount
+ or 0
+ )
+
+ if examinee_ids_to_check:
+ for eid in examinee_ids_to_check:
+ ref = db.scalar(
+ select(func.count())
+ .select_from(ExaminationExamExaminee)
+ .where(ExaminationExamExaminee.examinee_id == eid)
+ )
+ if ref == 0:
+ db.execute(delete(ExaminationExaminee).where(ExaminationExaminee.id == eid))
+ counts["import_batch"] = db.execute(delete(ImportBatch).where(ImportBatch.id == batch_id)).rowcount
+
+ db.commit()
+ return counts
+
+
+def parse_batch_exam_ids(batch: ImportBatch) -> list[int]:
+ if batch.exam_ids_json:
+ try:
+ raw = json.loads(batch.exam_ids_json)
+ if isinstance(raw, list) and raw:
+ return [int(x) for x in raw]
+ except (json.JSONDecodeError, TypeError, ValueError):
+ pass
+ if batch.exam_id is not None:
+ return [int(batch.exam_id)]
+ return []
--
Gitblit v1.8.0