doc-forge-reds/backend/app/api/generation_points.py

125 lines
4.3 KiB
Python

import uuid
from fastapi import APIRouter, Depends, HTTPException, UploadFile, File, Form, Query
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_point import GenerationPoint
from app.models.template import Template
from app.models.ai_model import AIModel
from app.schemas.generation_point import (
GenerationPointCreate,
GenerationPointUpdate,
GenerationPointResponse,
BatchOrderUpdate,
)
from app.services.file_storage import save_upload, REF_FILES_DIR
router = APIRouter(prefix="/generation-points", tags=["生成点管理"])
@router.post("", response_model=GenerationPointResponse)
async def create_generation_point(
template_id: uuid.UUID = Form(...),
position: str = Form(...),
prompt: str = Form(...),
model_id: uuid.UUID | None = Form(None),
order: int = Form(0),
selected_text: str | None = Form(None),
ref_file: UploadFile | None = File(None),
db: AsyncSession = Depends(get_db),
):
import json
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="模板不存在")
if model_id:
model_result = await db.execute(select(AIModel).where(AIModel.id == model_id))
if not model_result.scalar_one_or_none():
raise HTTPException(status_code=404, detail="模型不存在")
ref_file_path = None
if ref_file and ref_file.filename:
validate_file_extension(ref_file.filename)
ref_file_path = await save_upload(ref_file, REF_FILES_DIR)
point = GenerationPoint(
template_id=template_id,
position=json.loads(position),
prompt=prompt,
model_id=model_id,
order=order,
ref_file_path=ref_file_path,
selected_text=selected_text,
)
db.add(point)
await db.flush()
await db.refresh(point)
return point
@router.get("", response_model=list[GenerationPointResponse])
async def list_generation_points(
template_id: uuid.UUID = Query(...),
db: AsyncSession = Depends(get_db),
):
query = (
select(GenerationPoint)
.where(GenerationPoint.template_id == template_id)
.order_by(GenerationPoint.order.asc(), GenerationPoint.created_at.asc())
)
result = await db.execute(query)
return result.scalars().all()
@router.get("/{point_id}", response_model=GenerationPointResponse)
async def get_generation_point(point_id: str, db: AsyncSession = Depends(get_db)):
result = await db.execute(select(GenerationPoint).where(GenerationPoint.id == point_id))
point = result.scalar_one_or_none()
if not point:
raise HTTPException(status_code=404, detail="生成点不存在")
return point
@router.put("/{point_id}", response_model=GenerationPointResponse)
async def update_generation_point(
point_id: str,
data: GenerationPointUpdate,
db: AsyncSession = Depends(get_db),
):
result = await db.execute(select(GenerationPoint).where(GenerationPoint.id == point_id))
point = result.scalar_one_or_none()
if not point:
raise HTTPException(status_code=404, detail="生成点不存在")
update_data = data.model_dump(exclude_unset=True)
for key, value in update_data.items():
setattr(point, key, value)
await db.flush()
await db.refresh(point)
return point
@router.delete("/{point_id}")
async def delete_generation_point(point_id: str, db: AsyncSession = Depends(get_db)):
result = await db.execute(select(GenerationPoint).where(GenerationPoint.id == point_id))
point = result.scalar_one_or_none()
if not point:
raise HTTPException(status_code=404, detail="生成点不存在")
await db.delete(point)
return {"detail": "删除成功"}
@router.post("/batch-order")
async def batch_update_order(data: BatchOrderUpdate, db: AsyncSession = Depends(get_db)):
for item in data.points:
result = await db.execute(select(GenerationPoint).where(GenerationPoint.id == item["id"]))
point = result.scalar_one_or_none()
if point:
point.order = item["order"]
await db.flush()
return {"detail": "排序更新成功"}