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 []
|