252 lines
11 KiB
Python
252 lines
11 KiB
Python
from io import BytesIO
|
||
|
||
from fastapi import APIRouter, Depends, File, HTTPException, Query, UploadFile, status
|
||
from fastapi.responses import StreamingResponse
|
||
from openpyxl import Workbook, load_workbook
|
||
from openpyxl.comments import Comment
|
||
from openpyxl.styles import Alignment, Font, PatternFill
|
||
from sqlalchemy import func
|
||
from sqlalchemy import or_
|
||
from sqlalchemy.orm import Session
|
||
|
||
from app.api.deps import require_roles
|
||
from app.db.session import get_db
|
||
from app.models import AdminUser, Certificate, Learner, LearnerNameHistory
|
||
from app.schemas.learner import LearnerCreate, LearnerImportResult, LearnerOut, LearnerUpdate
|
||
from app.services.learner_identity import LearnerIdentityConflict, normalize_name, normalize_phone, resolve_learner
|
||
from app.services.logs import diff_values, log_action, mask_phone
|
||
|
||
router = APIRouter()
|
||
|
||
LEARNER_IMPORT_HEADERS = ["姓名", "手机号", "学员编号", "备注"]
|
||
|
||
|
||
@router.get("", response_model=list[LearnerOut])
|
||
def list_learners(
|
||
keyword: str | None = Query(default=None),
|
||
db: Session = Depends(get_db),
|
||
_: AdminUser = Depends(require_roles("system_admin", "certificate_admin", "readonly")),
|
||
) -> list[dict[str, object]]:
|
||
count_subquery = (
|
||
db.query(Certificate.learner_id.label("learner_id"), func.count(Certificate.id).label("certificate_count"))
|
||
.group_by(Certificate.learner_id)
|
||
.subquery()
|
||
)
|
||
query = (
|
||
db.query(Learner, func.coalesce(count_subquery.c.certificate_count, 0))
|
||
.outerjoin(count_subquery, count_subquery.c.learner_id == Learner.id)
|
||
.filter(Learner.status != "deleted")
|
||
.order_by(Learner.id.desc())
|
||
)
|
||
if keyword:
|
||
like = f"%{keyword}%"
|
||
query = query.filter(or_(Learner.current_name.like(like), Learner.phone.like(like), Learner.student_no.like(like)))
|
||
return [{**LearnerOut.model_validate(learner).model_dump(), "certificate_count": certificate_count} for learner, certificate_count in query.limit(100).all()]
|
||
|
||
|
||
@router.post("", response_model=LearnerOut, status_code=status.HTTP_201_CREATED)
|
||
def create_learner(
|
||
payload: LearnerCreate,
|
||
db: Session = Depends(get_db),
|
||
admin: AdminUser = Depends(require_roles("system_admin", "certificate_admin")),
|
||
) -> Learner:
|
||
try:
|
||
learner, created = resolve_learner(db, payload.current_name, payload.phone, source="manual", create_if_missing=True)
|
||
except (ValueError, LearnerIdentityConflict) as exc:
|
||
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail=str(exc)) from exc
|
||
assert learner is not None
|
||
learner.student_no = payload.student_no
|
||
learner.remark = payload.remark
|
||
action = "create_learner" if created else "update_learner"
|
||
log_action(db, admin, action, "learner", learner.id, {"name": learner.current_name, "phone": mask_phone(learner.phone), "status": learner.status})
|
||
db.commit()
|
||
db.refresh(learner)
|
||
return learner
|
||
|
||
|
||
@router.get("/import-template")
|
||
def download_learner_import_template(
|
||
_: AdminUser = Depends(require_roles("system_admin", "certificate_admin")),
|
||
) -> StreamingResponse:
|
||
workbook = build_learner_import_template()
|
||
stream = BytesIO()
|
||
workbook.save(stream)
|
||
stream.seek(0)
|
||
return StreamingResponse(
|
||
stream,
|
||
media_type="application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
|
||
headers={"Content-Disposition": 'attachment; filename="learner-import-template.xlsx"'},
|
||
)
|
||
|
||
|
||
@router.post("/import", response_model=LearnerImportResult)
|
||
def import_learners(
|
||
file: UploadFile = File(...),
|
||
db: Session = Depends(get_db),
|
||
admin: AdminUser = Depends(require_roles("system_admin", "certificate_admin")),
|
||
) -> dict[str, object]:
|
||
if not file.filename or not file.filename.lower().endswith(".xlsx"):
|
||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="仅支持 .xlsx 文件")
|
||
try:
|
||
workbook = load_workbook(file.file, read_only=True, data_only=True)
|
||
except Exception as exc:
|
||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Excel 文件无法读取") from exc
|
||
sheet = workbook.active
|
||
first_row = next(sheet.iter_rows(min_row=1, max_row=1), None)
|
||
headers = [cell.value for cell in first_row] if first_row else []
|
||
header_index = {name: idx for idx, name in enumerate(headers)}
|
||
missing = [name for name in LEARNER_IMPORT_HEADERS[:2] if name not in header_index]
|
||
if missing:
|
||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=f"缺少必填列:{'、'.join(missing)}")
|
||
|
||
result: dict[str, object] = {"total_rows": 0, "created_rows": 0, "updated_rows": 0, "unchanged_rows": 0, "failed_rows": 0, "errors": []}
|
||
seen_phones: dict[str, str] = {}
|
||
for row_no, row in enumerate(sheet.iter_rows(min_row=2, values_only=True), start=2):
|
||
if not any(row):
|
||
continue
|
||
result["total_rows"] += 1
|
||
values = {name: row[index] if index < len(row) else None for name, index in header_index.items()}
|
||
try:
|
||
name = normalize_name(str(values.get("姓名") or ""))
|
||
phone = normalize_phone(str(values.get("手机号") or ""))
|
||
if not name:
|
||
raise ValueError("姓名不能为空")
|
||
if phone in seen_phones and seen_phones[phone] != name:
|
||
raise LearnerIdentityConflict(f"文件内手机号重复且姓名不一致:{seen_phones[phone]} / {name}")
|
||
seen_phones[phone] = name
|
||
learner, created = resolve_learner(db, name, phone, source="learner_import", create_if_missing=True)
|
||
assert learner is not None
|
||
changed = False
|
||
student_no = _optional_cell(values.get("学员编号"))
|
||
remark = _optional_cell(values.get("备注"))
|
||
if student_no is not None and learner.student_no != student_no:
|
||
learner.student_no = student_no
|
||
changed = True
|
||
if remark is not None and learner.remark != remark:
|
||
learner.remark = remark
|
||
changed = True
|
||
if created:
|
||
result["created_rows"] += 1
|
||
elif changed:
|
||
result["updated_rows"] += 1
|
||
else:
|
||
result["unchanged_rows"] += 1
|
||
except (ValueError, LearnerIdentityConflict) as exc:
|
||
result["failed_rows"] += 1
|
||
result["errors"].append({"row_no": row_no, "message": str(exc)})
|
||
|
||
log_action(db, admin, "import_learners", "learner", detail={key: value for key, value in result.items() if key != "errors"})
|
||
db.commit()
|
||
return result
|
||
|
||
|
||
def build_learner_import_template() -> Workbook:
|
||
workbook = Workbook()
|
||
sheet = workbook.active
|
||
sheet.title = "学员导入模板"
|
||
sheet.append(LEARNER_IMPORT_HEADERS)
|
||
examples = {
|
||
"姓名": ("张三", "必填。填写学员真实姓名"),
|
||
"手机号": ("13800000000", "必填。11位手机号;系统中已存在时姓名必须一致"),
|
||
"学员编号": ("HY20260001", "选填。内部管理编号"),
|
||
"备注": ("2026年大本营学员", "选填。内部备注"),
|
||
}
|
||
for index, header in enumerate(LEARNER_IMPORT_HEADERS, start=1):
|
||
cell = sheet.cell(1, index)
|
||
cell.fill = PatternFill("solid", fgColor="208A87")
|
||
cell.font = Font(color="FFFFFF", bold=True)
|
||
cell.alignment = Alignment(horizontal="center", vertical="center")
|
||
cell.comment = Comment(examples[header][1], "证书管理系统")
|
||
sheet.column_dimensions[cell.column_letter].width = 18 if header != "备注" else 32
|
||
sheet.column_dimensions["B"].number_format = "@"
|
||
instruction = workbook.create_sheet("填写说明")
|
||
instruction.append(["字段", "是否必填", "示例", "填写说明"])
|
||
for header in LEARNER_IMPORT_HEADERS:
|
||
instruction.append([header, "是" if header in LEARNER_IMPORT_HEADERS[:2] else "否", examples[header][0], examples[header][1]])
|
||
instruction.column_dimensions["A"].width = 18
|
||
instruction.column_dimensions["B"].width = 14
|
||
instruction.column_dimensions["C"].width = 22
|
||
instruction.column_dimensions["D"].width = 52
|
||
sheet.freeze_panes = "A2"
|
||
return workbook
|
||
|
||
|
||
def _optional_cell(value: object) -> str | None:
|
||
text = str(value).strip() if value is not None else ""
|
||
return text or None
|
||
|
||
|
||
@router.get("/{learner_id}", response_model=LearnerOut)
|
||
def get_learner(
|
||
learner_id: int,
|
||
db: Session = Depends(get_db),
|
||
_: AdminUser = Depends(require_roles("system_admin", "certificate_admin", "readonly")),
|
||
) -> Learner:
|
||
learner = db.get(Learner, learner_id)
|
||
if not learner:
|
||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Learner not found")
|
||
return learner
|
||
|
||
|
||
@router.put("/{learner_id}", response_model=LearnerOut)
|
||
def update_learner(
|
||
learner_id: int,
|
||
payload: LearnerUpdate,
|
||
db: Session = Depends(get_db),
|
||
admin: AdminUser = Depends(require_roles("system_admin", "certificate_admin")),
|
||
) -> Learner:
|
||
learner = db.get(Learner, learner_id)
|
||
if not learner:
|
||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Learner not found")
|
||
try:
|
||
normalized_phone = normalize_phone(payload.phone)
|
||
normalized_name = normalize_name(payload.current_name)
|
||
except ValueError as exc:
|
||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)) from exc
|
||
duplicate = db.query(Learner).filter(Learner.phone == normalized_phone, Learner.id != learner_id).first()
|
||
if duplicate:
|
||
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="Phone already exists")
|
||
before = {
|
||
"phone": learner.phone,
|
||
"current_name": learner.current_name,
|
||
"student_no": learner.student_no,
|
||
"status": learner.status,
|
||
"remark": learner.remark,
|
||
}
|
||
if learner.current_name != normalized_name:
|
||
db.add(LearnerNameHistory(learner_id=learner.id, name=normalized_name, source="manual"))
|
||
learner.phone = normalized_phone
|
||
learner.current_name = normalized_name
|
||
learner.student_no = payload.student_no
|
||
learner.status = payload.status
|
||
learner.remark = payload.remark
|
||
after = {
|
||
"phone": learner.phone,
|
||
"current_name": learner.current_name,
|
||
"student_no": learner.student_no,
|
||
"status": learner.status,
|
||
"remark": learner.remark,
|
||
}
|
||
changes = diff_values(before, after, list(after.keys()))
|
||
if "phone" in changes:
|
||
changes["phone"] = {"before": mask_phone(changes["phone"]["before"]), "after": mask_phone(changes["phone"]["after"])}
|
||
log_action(db, admin, "update_learner", "learner", learner.id, {"name": learner.current_name, "phone": mask_phone(learner.phone), "changes": changes})
|
||
db.commit()
|
||
db.refresh(learner)
|
||
return learner
|
||
|
||
|
||
@router.delete("/{learner_id}")
|
||
def delete_learner(
|
||
learner_id: int,
|
||
db: Session = Depends(get_db),
|
||
admin: AdminUser = Depends(require_roles("system_admin", "certificate_admin")),
|
||
) -> dict[str, bool]:
|
||
learner = db.get(Learner, learner_id)
|
||
if not learner:
|
||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Learner not found")
|
||
learner.status = "deleted"
|
||
log_action(db, admin, "delete_learner", "learner", learner.id, {"name": learner.current_name, "phone": mask_phone(learner.phone)})
|
||
db.commit()
|
||
return {"ok": True}
|