125 lines
4.3 KiB
Python
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": "排序更新成功"}
|