init: 初始化项目

This commit is contained in:
zwt13703
2026-07-08 20:02:29 +08:00
parent 22590ae7b8
commit 1bb84df4ca
98 changed files with 8535 additions and 2 deletions
+233
View File
@@ -0,0 +1,233 @@
import json
import os
import uuid
from datetime import datetime, timezone
from fastapi import APIRouter, Depends, HTTPException, Query, UploadFile, File
from fastapi.responses import Response
from typing import List
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select
from app.core.database import get_db
from app.core.security_middleware import validate_file_extension
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, TaskDetailResponse
from app.services.document_processor import docx_to_pdf_bytes
from app.services.file_storage import get_file_content, save_upload, REF_FILES_DIR
from app.services.ai_adapter import call_ai_model
from app.services.ref_parser import parse_reference_file, parse_reference_files
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}")
async def get_task_status(task_id: str, detail: bool = Query(False), 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 detail:
tpl_result = await db.execute(select(Template).where(Template.id == task.template_id))
tpl = tpl_result.scalar_one_or_none()
points_result = await db.execute(
select(GenerationPoint).where(GenerationPoint.template_id == task.template_id).order_by(GenerationPoint.order.asc())
)
points = points_result.scalars().all()
return TaskDetailResponse(
id=task.id,
template_id=task.template_id,
template_name=tpl.name if tpl else "",
status=task.status,
result_file_path=task.result_file_path,
error_msg=task.error_msg,
created_at=task.created_at,
finished_at=task.finished_at,
points=[{
"id": str(p.id),
"prompt": p.prompt,
"position": p.position,
"model_id": str(p.model_id) if p.model_id else None,
"ref_file_path": p.ref_file_path,
"selected_text": p.selected_text,
"need_ref_file": p.need_ref_file,
"remark": p.remark,
"order": p.order,
} for p in points],
)
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,
ref_files: List[UploadFile] = File(default=[]),
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="生成点不存在")
tmp_paths = []
if point.need_ref_file:
if not ref_files and not point.ref_file_path:
raise HTTPException(status_code=400, detail="此生成点需要上传参考文件")
for f in ref_files:
if f.filename:
validate_file_extension(f.filename)
path = await save_upload(f, REF_FILES_DIR)
tmp_paths.append(path)
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 = ""
if tmp_paths:
ref_content = await parse_reference_files(json.dumps(tmp_paths))
elif point.ref_file_path:
ref_content = await parse_reference_files(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)}")
@router.post("/generation-points/{point_id}/upload-ref")
async def upload_ref_files(
point_id: str,
ref_files: List[UploadFile] = File(...),
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="生成点不存在")
paths = []
for f in ref_files:
if f.filename:
validate_file_extension(f.filename)
path = await save_upload(f, REF_FILES_DIR)
paths.append(path)
existing = []
if point.ref_file_path:
try:
existing = json.loads(point.ref_file_path)
if not isinstance(existing, list):
existing = [point.ref_file_path]
except json.JSONDecodeError:
existing = [point.ref_file_path] if point.ref_file_path else []
all_paths = existing + paths
point.ref_file_path = json.dumps(all_paths)
await db.commit()
return {"success": True, "count": len(all_paths)}