Files
doc-forge/backend/routers/generate.py
T
2026-07-02 15:00:11 +08:00

219 lines
7.3 KiB
Python

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})