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(SQL 文件中的试卷 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 []