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: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="导入批次不存在") failed_count = db.query(ImportBatchRow).filter(ImportBatchRow.batch_id == batch.id, ImportBatchRow.status == "failed").count() if not failed_count: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="错误报告不存在") batch.error_report_path = str(write_error_report(db, batch.id)) db.commit() 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 if failed_rows: db.flush() batch.error_report_path = str(write_error_report(db, batch.id)) 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 = normalize_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: db.flush() 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 normalize_row_data(row_data: dict[str, object]) -> dict[str, object]: normalized: dict[str, object] = {} for key, value in row_data.items(): if isinstance(value, datetime): normalized[key] = value.date().isoformat() elif isinstance(value, date): normalized[key] = value.isoformat() else: normalized[key] = value return normalized 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() try: return date.fromisoformat(text) except ValueError: pass try: return datetime.fromisoformat(text.replace("Z", "+00:00")).date() except ValueError: pass 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: db.flush() batch = db.get(ImportBatch, batch_id) template = get_certificate_template(batch.template_code) if batch else get_certificate_template("classic") headers = template_headers(template) workbook = Workbook() sheet = workbook.active sheet.title = "错误报告" sheet.append(["行号", "错误原因", *headers]) rows = db.query(ImportBatchRow).filter(ImportBatchRow.batch_id == batch_id, ImportBatchRow.status == "failed").all() for row in rows: raw_data = json.loads(row.raw_json or "{}") sheet.append([row.row_no, row.error_message or "未知错误", *(report_cell_value(header, raw_data.get(header)) for header in headers)]) sheet.freeze_panes = "A2" sheet.auto_filter.ref = f"A1:{sheet.cell(1, len(headers) + 2).coordinate}" sheet.column_dimensions["A"].width = 10 sheet.column_dimensions["B"].width = 48 for cell in sheet[1]: cell.fill = PatternFill("solid", fgColor="C0392B") cell.font = Font(color="FFFFFF", bold=True) cell.alignment = Alignment(horizontal="center", vertical="center") for row in sheet.iter_rows(min_row=2): row[1].alignment = Alignment(wrap_text=True, vertical="top") report_path = data_path("error-reports") / f"import-errors-{batch_id}.xlsx" workbook.save(report_path) return report_path def report_cell_value(header: str, value: object) -> object: if header not in {COL_COURSE_START_DATE, COL_COURSE_END_DATE, COL_ISSUE_DATE} or value in (None, ""): return value try: return parse_date(value, header).isoformat() except ValueError: return value 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"