285 lines
9.8 KiB
Python
285 lines
9.8 KiB
Python
import json
|
|
import time
|
|
import os
|
|
import uuid
|
|
from datetime import datetime
|
|
|
|
from fastapi import APIRouter, Depends, File, HTTPException, Query, UploadFile
|
|
from sqlalchemy import func, select
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from config import settings
|
|
from database import get_db
|
|
from models.ai_model import AiModel
|
|
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
|
|
from services.ai_service import call_ai
|
|
from services.minio_client import upload_bytes
|
|
|
|
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}
|
|
|
|
|
|
async def _get_effective_model(db: AsyncSession, paragraph: Paragraph) -> AiModel | None:
|
|
if paragraph.model_id:
|
|
model = await db.get(AiModel, paragraph.model_id)
|
|
if model is not None and model.status == "enabled":
|
|
return model
|
|
result = await db.execute(
|
|
select(AiModel).where(AiModel.status == "enabled").order_by(AiModel.id.asc()).limit(1)
|
|
)
|
|
return result.scalars().first()
|
|
|
|
|
|
@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="段落不存在")
|
|
|
|
model = None
|
|
if body.model_id:
|
|
model = await db.get(AiModel, body.model_id)
|
|
elif paragraph.model_id:
|
|
model = await db.get(AiModel, paragraph.model_id)
|
|
|
|
if model is None or model.status != "enabled":
|
|
content = _build_mock_content(paragraph)
|
|
message = "当前未找到可用模型,返回本地模拟生成结果。"
|
|
else:
|
|
result = await call_ai(paragraph, model)
|
|
content = result.content
|
|
message = f"已通过模型 {result.used_model} 生成。"
|
|
return Response(
|
|
data={
|
|
"paragraph_id": paragraph.id,
|
|
"content": content,
|
|
"message": message,
|
|
}
|
|
)
|
|
|
|
|
|
@router.post("/upload")
|
|
async def upload_reference_file(file: UploadFile = File(...)):
|
|
if not file.filename:
|
|
raise HTTPException(status_code=400, detail="文件名不能为空")
|
|
|
|
ext = os.path.splitext(file.filename)[1].lower()
|
|
if ext not in settings.ALLOWED_EXTENSIONS:
|
|
raise HTTPException(status_code=400, detail="文件类型不支持")
|
|
|
|
content = await file.read()
|
|
if not content:
|
|
raise HTTPException(status_code=400, detail="上传文件不能为空")
|
|
if len(content) > settings.MAX_UPLOAD_SIZE:
|
|
raise HTTPException(status_code=400, detail="文件大小超过限制")
|
|
|
|
object_name = f"{datetime.now().strftime('%Y%m%d')}/{uuid.uuid4().hex}{ext}"
|
|
await asyncio.to_thread(
|
|
upload_bytes,
|
|
settings.MINIO_BUCKET_UPLOADS,
|
|
object_name,
|
|
content,
|
|
file.content_type or "application/octet-stream",
|
|
)
|
|
return Response(
|
|
data={
|
|
"file_name": file.filename,
|
|
"file_path": f"{settings.MINIO_BUCKET_UPLOADS}/{object_name}",
|
|
}
|
|
)
|
|
|
|
|
|
@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
|
|
failed_count = 0
|
|
for paragraph in paragraphs:
|
|
if paragraph.edit_mode == "manual":
|
|
content = {"content": [{"type": "text", "text": paragraph.content or "该段落为人工编辑模式。"}]}
|
|
status = "success"
|
|
duration = 0
|
|
error_message = ""
|
|
else:
|
|
start = time.perf_counter()
|
|
model = await _get_effective_model(db, paragraph)
|
|
try:
|
|
if model is None:
|
|
content = _build_mock_content(paragraph)
|
|
else:
|
|
result = await call_ai(paragraph, model)
|
|
content = result.content
|
|
status = "success"
|
|
error_message = ""
|
|
except Exception as error:
|
|
content = _build_mock_content(paragraph)
|
|
status = "failed"
|
|
error_message = str(error)
|
|
failed_count += 1
|
|
duration = round(time.perf_counter() - start, 4)
|
|
|
|
log = GenerationLog(
|
|
document_id=document.id,
|
|
paragraph_id=paragraph.id,
|
|
model_id=paragraph.model_id,
|
|
status=status,
|
|
content=json.dumps(content, ensure_ascii=False),
|
|
duration=duration,
|
|
error_msg=error_message,
|
|
)
|
|
db.add(log)
|
|
done_count += 1
|
|
|
|
document.para_count_done = done_count
|
|
document.status = "completed" if failed_count == 0 else "failed"
|
|
document.error = "" if failed_count == 0 else f"{failed_count} 个段落生成失败,已回退为模拟结果。"
|
|
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})
|