""" API 路由:所有 REST 端点 """ import asyncio import json import logging import time from pathlib import Path from typing import Optional import numpy as np from fastapi import APIRouter, HTTPException, BackgroundTasks, Query from fastapi.responses import HTMLResponse, StreamingResponse, FileResponse from pydantic import BaseModel from .state import app_state from ..config import OUTPUT_DIR, DIMENSION_FRAMEWORK, LEVEL_DESCRIPTIONS from ..engines.llm_engine import LLMEngine from ..engines.report_renderer import ReportRenderer logger = logging.getLogger(__name__) router = APIRouter() # ==================== Pydantic Models ==================== class GenerateRequest(BaseModel): school: str use_cache: bool = True skip_llm: bool = False class BatchGenerateRequest(BaseModel): schools: Optional[list[str]] = None # None = 全部 use_cache: bool = True skip_llm: bool = False # ==================== 学校相关 ==================== @router.get("/schools") async def list_schools(): """获取所有学校的摘要列表""" return { "schools": app_state.get_all_schools_summary(), "total": len(app_state.schools), } @router.get("/schools/{school_name}") async def get_school_detail(school_name: str): """获取单个学校的详细数据""" if school_name not in app_state.schools: raise HTTPException(404, f"学校 '{school_name}' 不存在") report_data = app_state.get_report_data(school_name) # JSON 安全序列化 def _safe(obj): if isinstance(obj, (np.integer,)): return int(obj) if isinstance(obj, (np.floating,)): return float(obj) if isinstance(obj, np.ndarray): return obj.tolist() raise TypeError(f"Type {type(obj)} not serializable") # 转为 JSON 安全格式 safe_data = json.loads(json.dumps(report_data, default=_safe, ensure_ascii=False)) return safe_data @router.get("/schools/{school_name}/dimensions") async def get_school_dimensions(school_name: str): """获取学校的维度得分摘要""" if school_name not in app_state.schools: raise HTTPException(404, f"学校 '{school_name}' 不存在") report_data = app_state.get_report_data(school_name) return { "school": school_name, "overall": report_data["overall"], "dimensions": report_data["dimensions"], "sub_dimensions": report_data["sub_dimensions"], } # ==================== 报告生成 ==================== # 全局生成状态追踪 _generation_status = {} @router.post("/reports/generate") async def generate_report(req: GenerateRequest, background_tasks: BackgroundTasks): """ 生成单校报告(异步后台任务) 返回任务 ID,前端通过 SSE 监听进度 """ school = req.school if school not in app_state.schools: raise HTTPException(404, f"学校 '{school}' 不存在") task_id = f"gen_{school}_{int(time.time())}" _generation_status[task_id] = { "status": "pending", "school": school, "progress": 0, "total": 29, "current_segment": "", "started_at": time.time(), } background_tasks.add_task( _do_generate, task_id, school, req.use_cache, req.skip_llm ) return {"task_id": task_id, "school": school, "status": "started"} def _do_generate(task_id: str, school: str, use_cache: bool, skip_llm: bool): """后台执行报告生成""" status = _generation_status[task_id] status["status"] = "running" try: start = time.time() # 获取报告数据 report_data = app_state.get_report_data(school) # LLM 生成 if skip_llm: llm_sections = {} status["progress"] = status["total"] else: llm_engine = LLMEngine() def progress_cb(completed, total, segment_id): status["progress"] = completed status["total"] = total status["current_segment"] = segment_id llm_engine.set_progress_callback(progress_cb) llm_sections = llm_engine.generate_report_segments( report_data, use_cache=use_cache ) # 渲染 HTML renderer = ReportRenderer() output_path = OUTPUT_DIR / f"{school}_报告.html" renderer.render_to_file(report_data, llm_sections, output_path) # 保存数据 JSON def _safe(obj): if isinstance(obj, (np.integer,)): return int(obj) if isinstance(obj, (np.floating,)): return float(obj) if isinstance(obj, np.ndarray): return obj.tolist() raise TypeError(f"Type {type(obj)} not serializable") json_path = OUTPUT_DIR / f"{school}_report_data.json" with open(json_path, "w", encoding="utf-8") as f: json.dump(report_data, f, ensure_ascii=False, indent=2, default=_safe) elapsed = time.time() - start status["status"] = "completed" status["elapsed"] = round(elapsed, 1) status["output_path"] = str(output_path) status["report_size"] = output_path.stat().st_size status["score"] = report_data["overall"]["score"] status["rank"] = report_data["overall"]["rank_in_district"] logger.info(f"✅ {school} 报告生成完成,耗时 {elapsed:.1f}s") except Exception as e: status["status"] = "failed" status["error"] = str(e) logger.error(f"❌ {school} 报告生成失败: {e}", exc_info=True) @router.get("/reports/generate/{task_id}/status") async def get_generation_status(task_id: str): """查询生成任务状态""" if task_id not in _generation_status: raise HTTPException(404, f"任务 '{task_id}' 不存在") return _generation_status[task_id] @router.get("/reports/generate/{task_id}/stream") async def stream_generation_progress(task_id: str): """SSE 实时进度流""" if task_id not in _generation_status: raise HTTPException(404, f"任务 '{task_id}' 不存在") async def event_generator(): while True: status = _generation_status.get(task_id, {}) data = json.dumps(status, ensure_ascii=False, default=str) yield f"data: {data}\n\n" if status.get("status") in ("completed", "failed"): break await asyncio.sleep(0.5) return StreamingResponse( event_generator(), media_type="text/event-stream", headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"}, ) # ==================== 批量生成 ==================== _batch_status = {} @router.post("/reports/batch") async def batch_generate(req: BatchGenerateRequest, background_tasks: BackgroundTasks): """批量生成所有学校报告""" schools = req.schools or app_state.schools invalid = [s for s in schools if s not in app_state.schools] if invalid: raise HTTPException(400, f"无效学校: {invalid}") batch_id = f"batch_{int(time.time())}" _batch_status[batch_id] = { "status": "pending", "schools": schools, "total": len(schools), "completed": 0, "results": [], "started_at": time.time(), } background_tasks.add_task( _do_batch_generate, batch_id, schools, req.use_cache, req.skip_llm ) return {"batch_id": batch_id, "schools": schools, "total": len(schools)} def _do_batch_generate(batch_id: str, schools: list, use_cache: bool, skip_llm: bool): """后台执行批量生成""" status = _batch_status[batch_id] status["status"] = "running" renderer = ReportRenderer() llm_engine = None if skip_llm else LLMEngine() for idx, school in enumerate(schools): school_start = time.time() status["current_school"] = school status["current_index"] = idx try: report_data = app_state.get_report_data(school) if skip_llm: llm_sections = {} else: llm_sections = llm_engine.generate_report_segments( report_data, use_cache=use_cache ) output_path = OUTPUT_DIR / f"{school}_报告.html" renderer.render_to_file(report_data, llm_sections, output_path) # 保存数据 JSON def _safe(obj): if isinstance(obj, (np.integer,)): return int(obj) if isinstance(obj, (np.floating,)): return float(obj) if isinstance(obj, np.ndarray): return obj.tolist() raise TypeError(f"Type {type(obj)} not serializable") json_path = OUTPUT_DIR / f"{school}_report_data.json" with open(json_path, "w", encoding="utf-8") as f: json.dump(report_data, f, ensure_ascii=False, indent=2, default=_safe) elapsed = time.time() - school_start status["results"].append({ "school": school, "status": "success", "score": report_data["overall"]["score"], "rank": report_data["overall"]["rank_in_district"], "cluster": report_data["overall"]["cluster"], "time": round(elapsed, 1), }) except Exception as e: elapsed = time.time() - school_start status["results"].append({ "school": school, "status": "failed", "error": str(e), "time": round(elapsed, 1), }) logger.error(f"❌ 批量生成 {school} 失败: {e}", exc_info=True) status["completed"] = idx + 1 status["status"] = "completed" status["elapsed"] = round(time.time() - status["started_at"], 1) # 保存批量汇总 JSON summary_path = OUTPUT_DIR / "batch_summary.json" with open(summary_path, "w", encoding="utf-8") as f: json.dump({ "generated_at": time.strftime("%Y-%m-%d %H:%M:%S"), "total_schools": len(schools), "success_count": sum(1 for r in status["results"] if r["status"] == "success"), "total_time": status["elapsed"], "results": status["results"], }, f, ensure_ascii=False, indent=2) @router.get("/reports/batch/{batch_id}/status") async def get_batch_status(batch_id: str): """查询批量生成状态""" if batch_id not in _batch_status: raise HTTPException(404, f"批次 '{batch_id}' 不存在") return _batch_status[batch_id] @router.get("/reports/batch/{batch_id}/stream") async def stream_batch_progress(batch_id: str): """批量生成 SSE 进度流""" if batch_id not in _batch_status: raise HTTPException(404, f"批次 '{batch_id}' 不存在") async def event_generator(): while True: status = _batch_status.get(batch_id, {}) data = json.dumps(status, ensure_ascii=False, default=str) yield f"data: {data}\n\n" if status.get("status") in ("completed", "failed"): break await asyncio.sleep(1) return StreamingResponse( event_generator(), media_type="text/event-stream", headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"}, ) # ==================== 报告预览/下载 ==================== @router.get("/reports/{school_name}/preview") async def preview_report(school_name: str): """预览 HTML 报告(返回 HTML 内容)""" report_path = OUTPUT_DIR / f"{school_name}_报告.html" if not report_path.exists(): raise HTTPException(404, f"学校 '{school_name}' 的报告尚未生成") return HTMLResponse( content=report_path.read_text(encoding="utf-8"), media_type="text/html; charset=utf-8", ) @router.get("/reports/{school_name}/download") async def download_report(school_name: str): """下载 HTML 报告""" report_path = OUTPUT_DIR / f"{school_name}_报告.html" if not report_path.exists(): raise HTTPException(404, f"学校 '{school_name}' 的报告尚未生成") return FileResponse( path=str(report_path), filename=f"{school_name}_课程实施监测报告.html", media_type="text/html; charset=utf-8", ) @router.get("/reports/{school_name}/data") async def get_report_data(school_name: str): """获取报告的 JSON 数据""" json_path = OUTPUT_DIR / f"{school_name}_report_data.json" if not json_path.exists(): raise HTTPException(404, f"学校 '{school_name}' 的报告数据不存在") data = json.loads(json_path.read_text(encoding="utf-8")) return data # ==================== 系统配置 ==================== @router.get("/config/framework") async def get_framework(): """获取测评框架配置""" return { "dimensions": { name: { "sub_dimensions": info["sub_dimensions"], } for name, info in DIMENSION_FRAMEWORK.items() }, "level_descriptions": LEVEL_DESCRIPTIONS, } @router.get("/reports/summary") async def get_batch_summary(): """获取最近一次批量生成的汇总""" summary_path = OUTPUT_DIR / "batch_summary.json" if not summary_path.exists(): return {"message": "尚无批量生成记录"} return json.loads(summary_path.read_text(encoding="utf-8"))