import json import shutil from datetime import date, datetime from pathlib import Path from fastapi import APIRouter, Depends, File, Form, HTTPException, Query, UploadFile, status from fastapi.responses import FileResponse, StreamingResponse from openpyxl import Workbook, load_workbook from openpyxl.comments import Comment from openpyxl.styles import Alignment, Font, PatternFill from openpyxl.worksheet.datavalidation import DataValidation from sqlalchemy.orm import Session from app.api.deps import require_roles from app.core.paths import data_path from app.db.session import get_db from app.models import AdminUser, ImportBatch, ImportBatchRow, Learner, ProjectCourse from app.schemas.import_batch import ImportBatchOut from app.services.certificate_issuance import CertificateIssueData, DuplicateCertificate, issue_certificate from app.services.certificate_templates import CertificateTemplateDefinition, get_certificate_template from app.services.learner_identity import normalize_phone from app.services.logs import log_action router = APIRouter() COL_NAME = "姓名" COL_PHONE = "手机号" COL_PROJECT = "项目代码" COL_COURSE_NAME = "课程名称" COL_STAGE_NAME = "阶段名称" COL_COURSE_START_DATE = "课程开始日期" COL_COURSE_END_DATE = "课程结束日期" COL_ISSUE_DATE = "发证日期" COMMON_HEADERS = [COL_NAME, COL_PHONE, COL_PROJECT] FIELD_COLUMNS = { "course_name": COL_COURSE_NAME, "stage_name": COL_STAGE_NAME, "course_start_date": COL_COURSE_START_DATE, "course_end_date": COL_COURSE_END_DATE, "issue_date": COL_ISSUE_DATE, } TEMPLATE_HEADERS = COMMON_HEADERS + [ COL_COURSE_NAME, COL_STAGE_NAME, COL_COURSE_START_DATE, COL_COURSE_END_DATE, COL_ISSUE_DATE, ] REQUIRED_HEADERS = TEMPLATE_HEADERS @router.get("/template") def download_template( template_code: str = Query(default="classic"), _: AdminUser = Depends(require_roles("system_admin", "certificate_admin")), ) -> StreamingResponse: try: template = get_certificate_template(template_code) except (ValueError, FileNotFoundError) as exc: raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)) from exc workbook = build_import_template_workbook(template.code) stream_path = data_path("exports") / f"certificate-import-{template.code}.xlsx" workbook.save(stream_path) return StreamingResponse( stream_path.open("rb"), media_type="application/vnd.openxmlformats-officedocument.spreadsheetml.sheet", headers={"Content-Disposition": f'attachment; filename="certificate-import-{template.code}.xlsx"'}, ) def build_import_template_workbook(template_code: str = "classic") -> Workbook: template = get_certificate_template(template_code) headers = template_headers(template) workbook = Workbook() sheet = workbook.active sheet.title = "证书导入模板" sheet.append(headers) _format_template_sheet(sheet, headers, template) _add_template_instructions(workbook, template) return workbook @router.post("", response_model=ImportBatchOut, status_code=status.HTTP_201_CREATED) def upload_import_file( template_code: str = Form(...), file: UploadFile = File(...), db: Session = Depends(get_db), admin: AdminUser = Depends(require_roles("system_admin", "certificate_admin")), ) -> ImportBatch: if not file.filename or not file.filename.lower().endswith(".xlsx"): raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="仅支持 .xlsx 文件") try: template = get_certificate_template(template_code) except (ValueError, FileNotFoundError) as exc: raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)) from exc upload_path = data_path("uploads") / f"{datetime.now():%Y%m%d%H%M%S}-{Path(file.filename).name}" with upload_path.open("wb") as target: shutil.copyfileobj(file.file, target) batch = ImportBatch( filename=file.filename, file_path=str(upload_path), template_code=template.code, created_by=admin.id, ) db.add(batch) db.flush() validate_batch(db, batch, upload_path) log_action( db, admin, "upload_import_file", "import_batch", batch.id, {"filename": file.filename, "template_code": template.code, "total_rows": batch.total_rows, "valid_rows": batch.valid_rows, "failed_rows": batch.failed_rows, "status": batch.status}, ) db.commit() db.refresh(batch) return batch @router.get("", response_model=list[ImportBatchOut]) def list_import_batches( db: Session = Depends(get_db), _: AdminUser = Depends(require_roles("system_admin", "certificate_admin", "readonly")), ) -> list[ImportBatch]: return db.query(ImportBatch).order_by(ImportBatch.id.desc()).limit(50).all() @router.get("/{batch_id}/error-report") def download_error_report( batch_id: int, db: Session = Depends(get_db), _: AdminUser = Depends(require_roles("system_admin", "certificate_admin")), ) -> FileResponse: batch = db.get(ImportBatch, batch_id) if not batch or not batch.error_report_path: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="错误报告不存在") return FileResponse(batch.error_report_path, media_type="application/vnd.openxmlformats-officedocument.spreadsheetml.sheet", filename=f"import-errors-{batch.id}.xlsx") @router.get("/{batch_id}/file") def download_source_file( batch_id: int, db: Session = Depends(get_db), _: AdminUser = Depends(require_roles("system_admin", "certificate_admin")), ) -> FileResponse: batch = db.get(ImportBatch, batch_id) if not batch or not Path(batch.file_path).exists(): raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="原始文件不存在") return FileResponse(batch.file_path, media_type="application/vnd.openxmlformats-officedocument.spreadsheetml.sheet", filename=batch.filename) @router.delete("/{batch_id}") def delete_import_batch( batch_id: int, db: Session = Depends(get_db), admin: AdminUser = Depends(require_roles("system_admin", "certificate_admin")), ) -> dict[str, bool]: batch = db.get(ImportBatch, batch_id) if not batch: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="导入批次不存在") for file_name in [batch.file_path, batch.error_report_path]: if file_name and Path(file_name).exists(): Path(file_name).unlink() db.query(ImportBatchRow).filter(ImportBatchRow.batch_id == batch.id).delete() db.delete(batch) log_action(db, admin, "delete_import_batch", "import_batch", batch.id, {"filename": batch.filename, "status": batch.status, "total_rows": batch.total_rows}) db.commit() return {"ok": True} @router.post("/{batch_id}/confirm", response_model=ImportBatchOut) def confirm_import_batch( batch_id: int, db: Session = Depends(get_db), admin: AdminUser = Depends(require_roles("system_admin", "certificate_admin")), ) -> ImportBatch: batch = db.get(ImportBatch, batch_id) if not batch: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="导入批次不存在") if batch.status == "imported": return batch if batch.status != "validated": raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="当前批次不能确认导入") rows = db.query(ImportBatchRow).filter(ImportBatchRow.batch_id == batch.id, ImportBatchRow.status == "valid").all() imported_rows = skipped_rows = failed_rows = created_learners = 0 for row in rows: row_data = json.loads(row.raw_json or "{}") try: _, _, learner_created = issue_certificate( db, CertificateIssueData( learner_name=str(row_data[COL_NAME]), learner_phone=str(row_data[COL_PHONE]), project_code=str(row_data[COL_PROJECT]), template_code=batch.template_code, course_name=optional_text(row_data.get(COL_COURSE_NAME)), stage_name=optional_text(row_data.get(COL_STAGE_NAME)), course_start_date=parse_optional_date(row_data.get(COL_COURSE_START_DATE), COL_COURSE_START_DATE), course_end_date=parse_optional_date(row_data.get(COL_COURSE_END_DATE), COL_COURSE_END_DATE), issue_date=parse_issue_date(row_data[COL_ISSUE_DATE]), import_batch_id=batch.id, ), source="certificate_import", ) row.status = "imported" imported_rows += 1 created_learners += int(learner_created) except DuplicateCertificate as exc: row.status = "skipped" row.error_message = str(exc) skipped_rows += 1 except (ValueError, FileNotFoundError) as exc: row.status = "failed" row.error_message = str(exc) failed_rows += 1 batch.status = "imported" if imported_rows or skipped_rows else "failed" batch.failed_rows = (batch.failed_rows or 0) + failed_rows log_action( db, admin, "confirm_import_batch", "import_batch", batch.id, {"filename": batch.filename, "imported_rows": imported_rows, "skipped_rows": skipped_rows, "created_learners": created_learners, "failed_rows": failed_rows, "status": batch.status}, ) db.commit() db.refresh(batch) return batch def validate_batch(db: Session, batch: ImportBatch, upload_path: Path) -> None: template = get_certificate_template(batch.template_code) headers = template_headers(template) workbook = load_workbook(upload_path, read_only=True, data_only=True) sheet = workbook.active first_row = next(sheet.iter_rows(min_row=1, max_row=1), None) header = [cell.value for cell in first_row] if first_row else [] header_index = {name: idx for idx, name in enumerate(header)} missing = [name for name in required_headers(template) if name not in header_index] if missing: batch.status = "failed" batch.failed_rows = 1 db.add(ImportBatchRow(batch_id=batch.id, row_no=1, status="failed", error_message=f"缺少必填列:{'、'.join(missing)}")) return active_codes = {row[0] for row in db.query(ProjectCourse.code).filter(ProjectCourse.status == "active").all()} total = valid = failed = 0 for row_no, row in enumerate(sheet.iter_rows(min_row=2, values_only=True), start=2): if not any(row): continue total += 1 row_data = {name: row[header_index[name]] if name in header_index and header_index[name] < len(row) else None for name in headers} errors = row_errors(row_data, active_codes, template.code, db) row_status = "failed" if errors else "valid" failed += int(bool(errors)) valid += int(not errors) db.add( ImportBatchRow( batch_id=batch.id, row_no=row_no, status=row_status, error_message=";".join(errors) if errors else None, raw_json=json.dumps(row_data, ensure_ascii=False, default=str), ) ) batch.total_rows = total batch.valid_rows = valid batch.failed_rows = failed batch.status = "validated" if total else "failed" if failed: batch.error_report_path = str(write_error_report(db, batch.id)) def row_errors( row_data: dict[str, object], active_codes: set[str], template_code: str = "classic", db: Session | None = None, ) -> list[str]: template = get_certificate_template(template_code) errors: list[str] = [] for name in required_headers(template): if not optional_text(row_data.get(name)): errors.append(f"{name}不能为空") project_code = str(row_data.get(COL_PROJECT) or "").strip().upper() if project_code and project_code not in active_codes: errors.append("项目代码不存在或已停用") for field in template.fields: if field.field_type != "date" or field.key == "learner_name": continue column = FIELD_COLUMNS[field.key] if row_data.get(column) and not date_is_valid(row_data[column]): errors.append(f"{column}格式错误,请使用YYYY-MM-DD,例如{field.example}") if all(row_data.get(column) and date_is_valid(row_data[column]) for column in [COL_COURSE_START_DATE, COL_COURSE_END_DATE]): if parse_date(row_data[COL_COURSE_END_DATE], COL_COURSE_END_DATE) < parse_date(row_data[COL_COURSE_START_DATE], COL_COURSE_START_DATE): errors.append("课程结束日期不能早于课程开始日期") if row_data.get(COL_PHONE): try: phone = normalize_phone(str(row_data[COL_PHONE])) if db: learner = db.query(Learner).filter(Learner.phone == phone, Learner.status != "deleted").first() input_name = str(row_data.get(COL_NAME) or "").strip() if learner and learner.current_name.strip() != input_name: errors.append(f"手机号已属于学员“{learner.current_name}”,姓名不一致") except ValueError as exc: errors.append(str(exc)) return errors def template_headers(template: CertificateTemplateDefinition) -> list[str]: return COMMON_HEADERS + [FIELD_COLUMNS[field.key] for field in template.fields if field.key != "learner_name"] def required_headers(template: CertificateTemplateDefinition) -> list[str]: return COMMON_HEADERS + [FIELD_COLUMNS[field.key] for field in template.fields if field.key != "learner_name" and field.required] def optional_text(value: object) -> str | None: if value is None: return None text = str(value).strip() return text or None def parse_issue_date(value: object) -> date: return parse_date(value, COL_ISSUE_DATE) def parse_optional_date(value: object, field_name: str) -> date | None: return parse_date(value, field_name) if value not in (None, "") else None def parse_date(value: object, field_name: str = "日期") -> date: if isinstance(value, datetime): return value.date() if isinstance(value, date): return value text = str(value).strip() for fmt in ["%Y-%m-%d", "%Y/%m/%d", "%Y.%m.%d"]: try: return datetime.strptime(text, fmt).date() except ValueError: continue raise ValueError(f"{field_name}格式错误,请使用YYYY-MM-DD:{text}") def date_is_valid(value: object) -> bool: try: parse_date(value) return True except ValueError: return False def write_error_report(db: Session, batch_id: int) -> Path: workbook = Workbook() sheet = workbook.active sheet.title = "错误报告" sheet.append(["行号", "错误原因", "原始数据"]) rows = db.query(ImportBatchRow).filter(ImportBatchRow.batch_id == batch_id, ImportBatchRow.status == "failed").all() for row in rows: sheet.append([row.row_no, row.error_message, row.raw_json]) report_path = data_path("error-reports") / f"import-errors-{batch_id}.xlsx" workbook.save(report_path) return report_path def _format_template_sheet(sheet, headers: list[str], template: CertificateTemplateDefinition) -> None: header_fill = PatternFill("solid", fgColor="208A87") for cell in sheet[1]: cell.fill = header_fill cell.font = Font(color="FFFFFF", bold=True) cell.alignment = Alignment(horizontal="center", vertical="center") sheet.freeze_panes = "A2" sheet.auto_filter.ref = f"A1:{sheet.cell(1, len(headers)).coordinate}" widths = {COL_NAME: 16, COL_PHONE: 18, COL_PROJECT: 16, COL_COURSE_NAME: 28, COL_STAGE_NAME: 18, COL_COURSE_START_DATE: 18, COL_COURSE_END_DATE: 18, COL_ISSUE_DATE: 18} sheet.column_dimensions["B"].number_format = "@" for index, header in enumerate(headers, start=1): column_letter = sheet.cell(1, index).column_letter sheet.column_dimensions[column_letter].width = widths[header] field = next((item for item in template.fields if item.label == header), None) required = header in required_headers(template) sheet.cell(1, index).comment = Comment(f"{'必填' if required else '选填'}。{field.description if field else '用于识别和归档数据'}", "证书管理系统") if field and field.field_type == "date": sheet.column_dimensions[column_letter].number_format = "yyyy-mm-dd" validation = DataValidation(type="date", operator="between", formula1="DATE(2000,1,1)", formula2="DATE(2100,12,31)", allow_blank=not field.required) validation.promptTitle = "日期格式" validation.prompt = f"请按 YYYY-MM-DD 填写,例如 {field.example}" validation.errorTitle = "日期格式错误" validation.error = "请填写有效日期" validation.errorStyle = "stop" validation.showInputMessage = True validation.showErrorMessage = True sheet.add_data_validation(validation) validation.add(f"{column_letter}2:{column_letter}5000") def _add_template_instructions(workbook: Workbook, template: CertificateTemplateDefinition) -> None: sheet = workbook.create_sheet("填写说明") sheet.append(["字段", "是否必填", "格式或示例", "填写说明"]) instructions = [ (COL_NAME, True, "张三", "填写学员真实姓名;与手机号共同确认学员身份"), (COL_PHONE, True, "13800000000", "不存在时自动创建学员;已存在时姓名必须一致"), (COL_PROJECT, True, "DBY", "填写系统中已启用的项目代码"), ] instructions.extend((field.label, field.required, field.example, field.description) for field in template.fields if field.key != "learner_name") for label, required, example, description in instructions: sheet.append([label, "是" if required else "否", example, description]) note_row = len(instructions) + 3 sheet.cell(note_row, 1, "重要提示") sheet.cell(note_row, 2, f"本文件仅适用于“{template.name}”。不要修改第一行列名;日期统一使用 YYYY-MM-DD。") sheet.merge_cells(start_row=note_row, start_column=2, end_row=note_row, end_column=4) for cell in sheet[1]: cell.fill = PatternFill("solid", fgColor="208A87") cell.font = Font(color="FFFFFF", bold=True) cell.alignment = Alignment(horizontal="center") sheet.cell(note_row, 1).font = Font(color="C00000", bold=True) sheet.cell(note_row, 2).font = Font(color="C00000", bold=True) sheet.cell(note_row, 2).alignment = Alignment(wrap_text=True, vertical="center") sheet.row_dimensions[note_row].height = 34 sheet.column_dimensions["A"].width = 20 sheet.column_dimensions["B"].width = 16 sheet.column_dimensions["C"].width = 24 sheet.column_dimensions["D"].width = 58 sheet.freeze_panes = "A2"