import json import time from datetime import datetime from fastapi import APIRouter, Depends, HTTPException, Query from sqlalchemy import func, select from sqlalchemy.ext.asyncio import AsyncSession from database import get_db from models.document import Document from models.generation_log import GenerationLog from models.paragraph import Paragraph from models.template import Template from schemas.schemas import GenerateFullRequest, GenerateTestRequest, Response router = APIRouter() def _serialize_document(document: Document) -> dict: return { "id": document.id, "template_id": document.template_id, "name": document.name, "para_count_done": document.para_count_done, "para_count_total": document.para_count_total, "status": document.status, "file_path": document.file_path, "error": document.error, "created_at": document.created_at, "updated_at": document.updated_at, } def _build_mock_content(paragraph: Paragraph) -> dict: if paragraph.output_format == "table": return { "content": [ { "type": "table", "title": paragraph.title, "headers": ["字段", "内容"], "rows": [ ["段落标题", paragraph.title], ["生成说明", paragraph.prompt_text or "根据模板内容生成"], ], } ] } blocks = [ { "type": "text", "text": f"这是“{paragraph.title}”的示例生成内容,可用于前端联调与流程验证。" } ] if paragraph.content: blocks.append({"type": "text", "text": f"模板上下文:{paragraph.content[:200]}"}) if paragraph.need_prompt and paragraph.prompt_text: blocks.append({"type": "text", "text": f"预设提示词:{paragraph.prompt_text[:200]}"}) return {"content": blocks} @router.post("/test") async def generate_test(body: GenerateTestRequest, db: AsyncSession = Depends(get_db)): paragraph = await db.get(Paragraph, body.paragraph_id) if paragraph is None or paragraph.template_id != body.template_id: raise HTTPException(status_code=404, detail="段落不存在") content = _build_mock_content(paragraph) return Response( data={ "paragraph_id": paragraph.id, "content": content, "message": "当前返回本地模拟生成结果,便于前端联调。", } ) @router.post("/full") async def generate_full(body: GenerateFullRequest, db: AsyncSession = Depends(get_db)): template = await db.get(Template, body.template_id) if template is None: raise HTTPException(status_code=404, detail="模板不存在") result = await db.execute( select(Paragraph) .where(Paragraph.template_id == body.template_id) .order_by(Paragraph.sort_index.asc(), Paragraph.id.asc()) ) paragraphs = result.scalars().all() if not paragraphs: raise HTTPException(status_code=400, detail="模板下暂无可生成段落") document = Document( template_id=template.id, name=f"{template.name}-{datetime.now().strftime('%Y%m%d%H%M%S')}", para_count_done=0, para_count_total=len(paragraphs), status="generating", file_path="", error="", ) db.add(document) await db.flush() done_count = 0 for paragraph in paragraphs: if paragraph.edit_mode == "manual": content = {"content": [{"type": "text", "text": paragraph.content or "该段落为人工编辑模式。"}]} else: start = time.perf_counter() content = _build_mock_content(paragraph) duration = round(time.perf_counter() - start, 4) log = GenerationLog( document_id=document.id, paragraph_id=paragraph.id, model_id=paragraph.model_id, status="success", content=json.dumps(content, ensure_ascii=False), duration=duration, error_msg="", ) db.add(log) done_count += 1 continue log = GenerationLog( document_id=document.id, paragraph_id=paragraph.id, model_id=paragraph.model_id, status="success", content=json.dumps(content, ensure_ascii=False), duration=0, error_msg="", ) db.add(log) done_count += 1 document.para_count_done = done_count document.status = "completed" document.file_path = f"mock://document/{document.id}" await db.commit() await db.refresh(document) return Response(data=_serialize_document(document)) @router.get("/documents") async def list_documents( page: int = Query(1, ge=1), page_size: int = Query(20, ge=1, le=100), db: AsyncSession = Depends(get_db), ): total = (await db.execute(select(func.count(Document.id)))).scalar_one() result = await db.execute( select(Document) .order_by(Document.id.desc()) .offset((page - 1) * page_size) .limit(page_size) ) items = [_serialize_document(item) for item in result.scalars().all()] return Response(data={"items": items, "total": total, "page": page, "page_size": page_size}) @router.get("/documents/{document_id}") async def get_document(document_id: int, db: AsyncSession = Depends(get_db)): document = await db.get(Document, document_id) if document is None: raise HTTPException(status_code=404, detail="生成记录不存在") log_result = await db.execute( select(GenerationLog, Paragraph) .join(Paragraph, Paragraph.id == GenerationLog.paragraph_id) .where(GenerationLog.document_id == document_id) .order_by(Paragraph.sort_index.asc(), Paragraph.id.asc()) ) items = [] for log, paragraph in log_result.all(): items.append( { "id": log.id, "paragraph_id": paragraph.id, "title": paragraph.title, "sort_index": paragraph.sort_index, "status": log.status, "content": json.loads(log.content) if log.content else {"content": []}, } ) payload = _serialize_document(document) payload["logs"] = items return Response(data=payload) @router.post("/cancel/{document_id}") async def cancel_document(document_id: int, db: AsyncSession = Depends(get_db)): document = await db.get(Document, document_id) if document is None: raise HTTPException(status_code=404, detail="生成记录不存在") document.status = "cancelled" await db.commit() await db.refresh(document) return Response(data=_serialize_document(document)) @router.delete("/documents/{document_id}") async def delete_document(document_id: int, db: AsyncSession = Depends(get_db)): document = await db.get(Document, document_id) if document is None: raise HTTPException(status_code=404, detail="生成记录不存在") result = await db.execute(select(GenerationLog).where(GenerationLog.document_id == document_id)) for log in result.scalars().all(): await db.delete(log) await db.delete(document) await db.commit() return Response(data={"id": document_id})