import os import uuid from datetime import datetime, timezone from fastapi import APIRouter, Depends, HTTPException, Query from fastapi.responses import Response from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy import select from app.core.database import get_db from app.models.generation_task import GenerationTask from app.models.generation_point import GenerationPoint from app.models.template import Template from app.models.ai_model import AIModel from app.schemas.generation_task import TaskResponse, GenerateResponse, SingleTestRequest, SingleTestResponse from app.services.document_processor import docx_to_pdf_bytes from app.services.file_storage import get_file_content from app.services.ai_adapter import call_ai_model from app.services.ref_parser import parse_reference_file from app.tasks.generate import generate_document router = APIRouter() @router.post("/templates/{template_id}/generate", response_model=GenerateResponse) async def trigger_generation(template_id: str, db: AsyncSession = Depends(get_db)): template_result = await db.execute(select(Template).where(Template.id == template_id)) if not template_result.scalar_one_or_none(): raise HTTPException(status_code=404, detail="模板不存在") points_result = await db.execute( select(GenerationPoint).where(GenerationPoint.template_id == template_id) ) points = points_result.scalars().all() if not points: raise HTTPException(status_code=400, detail="模板没有生成点") task = GenerationTask( template_id=uuid.UUID(template_id), status="pending", created_at=datetime.now(timezone.utc), ) db.add(task) await db.flush() celery_task = generate_document.delay(str(task.id)) task.celery_task_id = celery_task.id await db.commit() return GenerateResponse(task_id=task.id, status="pending") @router.get("/tasks/{task_id}", response_model=TaskResponse) async def get_task_status(task_id: str, db: AsyncSession = Depends(get_db)): result = await db.execute(select(GenerationTask).where(GenerationTask.id == task_id)) task = result.scalar_one_or_none() if not task: raise HTTPException(status_code=404, detail="任务不存在") return task @router.get("/tasks/{task_id}/download") async def download_task_result( task_id: str, format: str = Query("docx", pattern="^(docx|pdf)$"), db: AsyncSession = Depends(get_db), ): result = await db.execute(select(GenerationTask).where(GenerationTask.id == task_id)) task = result.scalar_one_or_none() if not task: raise HTTPException(status_code=404, detail="任务不存在") if task.status != "done" or not task.result_file_path: raise HTTPException(status_code=400, detail="任务未完成或无结果文件") file_content = await get_file_content(task.result_file_path) if format == "pdf": file_content = docx_to_pdf_bytes(file_content) media_type = "application/pdf" filename = f"result_{task_id}.pdf" else: media_type = "application/vnd.openxmlformats-officedocument.wordprocessingml.document" filename = f"result_{task_id}.docx" return Response( content=file_content, media_type=media_type, headers={"Content-Disposition": f'attachment; filename="{filename}"'}, ) @router.get("/tasks") async def list_tasks( skip: int = Query(0, ge=0), limit: int = Query(20, ge=1, le=100), db: AsyncSession = Depends(get_db), ): query = select(GenerationTask).offset(skip).limit(limit).order_by(GenerationTask.created_at.desc()) result = await db.execute(query) return result.scalars().all() @router.post("/tasks/{task_id}/cancel") async def cancel_task(task_id: str, db: AsyncSession = Depends(get_db)): result = await db.execute(select(GenerationTask).where(GenerationTask.id == task_id)) task = result.scalar_one_or_none() if not task: raise HTTPException(status_code=404, detail="任务不存在") if task.status not in ("pending", "processing"): raise HTTPException(status_code=400, detail="任务无法取消") if task.celery_task_id: from app.tasks.celery_app import celery_app celery_app.control.revoke(task.celery_task_id, terminate=True) task.status = "failed" task.error_msg = "用户取消" task.finished_at = datetime.now(timezone.utc) await db.commit() return {"detail": "任务已取消"} @router.post("/generation-points/{point_id}/test", response_model=SingleTestResponse) async def test_single_point( point_id: str, db: AsyncSession = Depends(get_db), ): point_result = await db.execute(select(GenerationPoint).where(GenerationPoint.id == point_id)) point = point_result.scalar_one_or_none() if not point: raise HTTPException(status_code=404, detail="生成点不存在") model_config = {"provider": "custom", "endpoint": "", "api_key": "", "extra_params": {}} if point.model_id: model_result = await db.execute(select(AIModel).where(AIModel.id == point.model_id)) model = model_result.scalar_one_or_none() if model: model_config = { "provider": model.provider, "endpoint": model.endpoint, "api_key": model.api_key, "extra_params": model.extra_params, } ref_content = None if point.ref_file_path: ref_content = await parse_reference_file(point.ref_file_path) try: ai_result = await call_ai_model(model_config, point.prompt, ref_content) return SingleTestResponse(result=ai_result) except Exception as e: raise HTTPException(status_code=500, detail=f"AI 调用失败: {str(e)}")