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