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