Co-authored-by: factory-droid[bot] <138933559+factory-droid[bot]@users.noreply.github.com>
413 lines
13 KiB
Python
413 lines
13 KiB
Python
"""
|
|
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"))
|