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.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), 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: 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, ) 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": "排序更新成功"}