init: 初始化项目
This commit is contained in:
@@ -0,0 +1,128 @@
|
||||
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),
|
||||
need_ref_file: bool = Form(False),
|
||||
remark: 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,
|
||||
need_ref_file=need_ref_file,
|
||||
remark=remark,
|
||||
)
|
||||
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": "排序更新成功"}
|
||||
Reference in New Issue
Block a user