237 lines
7.6 KiB
Go
237 lines
7.6 KiB
Go
// 图片生成接口处理器,负责生图参数校验、参考图上传、历史查询和结果删除。
|
|
package handler
|
|
|
|
import (
|
|
"context"
|
|
"crypto/sha256"
|
|
"encoding/hex"
|
|
"errors"
|
|
"image"
|
|
"io"
|
|
"mime"
|
|
"net/http"
|
|
"path/filepath"
|
|
"strings"
|
|
|
|
mediakey "juhe-factory/api/internal/media"
|
|
"juhe-factory/api/internal/model"
|
|
"juhe-factory/api/internal/modules/productimage"
|
|
"juhe-factory/api/internal/service"
|
|
"juhe-factory/api/internal/storage"
|
|
|
|
"github.com/gin-gonic/gin"
|
|
"github.com/google/uuid"
|
|
"gorm.io/gorm"
|
|
)
|
|
|
|
// ProductImage 组合图片生成生图服务与对象存储依赖。
|
|
type ProductImage struct {
|
|
service *productimage.Service
|
|
cos *storage.COS
|
|
}
|
|
|
|
// NewProductImage 创建图片生成 HTTP 处理器。
|
|
func NewProductImage(service *productimage.Service, cos *storage.COS) *ProductImage {
|
|
return &ProductImage{service: service, cos: cos}
|
|
}
|
|
|
|
// Generate 接收当前页面的临时配置和参考图,并创建图片生成任务。
|
|
func (h *ProductImage) Generate(c *gin.Context) {
|
|
if h.cos == nil {
|
|
fail(c, http.StatusServiceUnavailable, "cos_not_configured", "对象存储未配置")
|
|
return
|
|
}
|
|
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, 48<<20)
|
|
if err := c.Request.ParseMultipartForm(48 << 20); err != nil {
|
|
fail(c, http.StatusBadRequest, "invalid_request", "生图请求无效")
|
|
return
|
|
}
|
|
if c.Request.MultipartForm != nil {
|
|
defer c.Request.MultipartForm.RemoveAll()
|
|
}
|
|
modelID, err := service.ParseUUID(c.PostForm("model_id"), "图片模型")
|
|
if err != nil {
|
|
fail(c, http.StatusBadRequest, "invalid_request", err.Error())
|
|
return
|
|
}
|
|
ratio := strings.TrimSpace(c.PostForm("aspect_ratio"))
|
|
if ratio == "" {
|
|
ratio = "9:16"
|
|
}
|
|
resolution := strings.TrimSpace(c.PostForm("resolution"))
|
|
if resolution == "" {
|
|
resolution = "1k"
|
|
}
|
|
userID := currentWebUser(c).ID
|
|
files := c.Request.MultipartForm.File["references"]
|
|
if len(files) > 4 {
|
|
fail(c, http.StatusBadRequest, "invalid_request", "最多上传 4 张参考图")
|
|
return
|
|
}
|
|
references := make([]productimage.ReferenceUpload, 0, len(files))
|
|
objectKeys := make([]string, 0, len(files))
|
|
for index, header := range files {
|
|
file, openErr := header.Open()
|
|
if openErr != nil {
|
|
if !h.cleanupUploadedObjects(c, objectKeys) {
|
|
return
|
|
}
|
|
fail(c, http.StatusBadRequest, "invalid_request", "参考图无法读取")
|
|
return
|
|
}
|
|
contentType := cleanProductImageContentType(header.Header.Get("Content-Type"))
|
|
if contentType == "" {
|
|
contentType = cleanProductImageContentType(mime.TypeByExtension(strings.ToLower(filepath.Ext(header.Filename))))
|
|
}
|
|
allowed := map[string]bool{"image/jpeg": true, "image/png": true, "image/webp": true}
|
|
if !allowed[contentType] || header.Size <= 0 || header.Size > h.cos.MaxImageBytes() {
|
|
file.Close()
|
|
if !h.cleanupUploadedObjects(c, objectKeys) {
|
|
return
|
|
}
|
|
fail(c, http.StatusBadRequest, "invalid_request", "参考图格式或大小不符合要求")
|
|
return
|
|
}
|
|
config, _, decodeErr := image.DecodeConfig(file)
|
|
if decodeErr != nil || config.Width <= 0 || config.Height <= 0 {
|
|
file.Close()
|
|
if !h.cleanupUploadedObjects(c, objectKeys) {
|
|
return
|
|
}
|
|
fail(c, http.StatusBadRequest, "invalid_request", "参考图内容无法解析")
|
|
return
|
|
}
|
|
if _, err = file.Seek(0, io.SeekStart); err != nil {
|
|
file.Close()
|
|
if !h.cleanupUploadedObjects(c, objectKeys) {
|
|
return
|
|
}
|
|
h.respondError(c, err)
|
|
return
|
|
}
|
|
mediaID := uuid.New()
|
|
key := mediakey.ProductImageReference(userID, mediaID, header.Filename, contentType)
|
|
hash := sha256.New()
|
|
if _, err = io.Copy(hash, file); err != nil {
|
|
file.Close()
|
|
if !h.cleanupUploadedObjects(c, objectKeys) {
|
|
return
|
|
}
|
|
h.respondError(c, err)
|
|
return
|
|
}
|
|
if _, err = file.Seek(0, io.SeekStart); err != nil {
|
|
file.Close()
|
|
if !h.cleanupUploadedObjects(c, objectKeys) {
|
|
return
|
|
}
|
|
h.respondError(c, err)
|
|
return
|
|
}
|
|
url, putErr := h.cos.Put(c, key, contentType, file, header.Size)
|
|
file.Close()
|
|
if putErr != nil {
|
|
if !h.cleanupUploadedObjects(c, objectKeys) {
|
|
return
|
|
}
|
|
h.respondError(c, putErr)
|
|
return
|
|
}
|
|
objectKeys = append(objectKeys, key)
|
|
name := strings.TrimSpace(filepath.Base(header.Filename))
|
|
if name == "" {
|
|
name = "参考图" + string(rune('1'+index))
|
|
}
|
|
owner := userID
|
|
references = append(references, productimage.ReferenceUpload{Asset: &model.MediaAsset{ID: mediaID, OwnerUserID: &owner, StorageProvider: "cos", ObjectKey: key, PublicURL: url, OriginalName: name, DisplayName: "图片生成参考图", MimeType: contentType, SizeBytes: header.Size, SHA256: hex.EncodeToString(hash.Sum(nil)), Width: &config.Width, Height: &config.Height}, Name: name})
|
|
}
|
|
task, err := h.service.QueueGeneration(userID, productimage.GenerateInput{Prompt: c.PostForm("prompt"), ModelID: modelID, AspectRatio: ratio, Resolution: resolution, References: references})
|
|
if err != nil {
|
|
if !h.cleanupUploadedObjects(c, objectKeys) {
|
|
return
|
|
}
|
|
h.respondError(c, err)
|
|
return
|
|
}
|
|
c.JSON(http.StatusAccepted, gin.H{"data": gin.H{"task_id": task.ID}})
|
|
}
|
|
|
|
// ListGenerations 返回当前用户最近的图片生成历史。
|
|
func (h *ProductImage) ListGenerations(c *gin.Context) {
|
|
items, err := h.service.Generations(currentWebUser(c).ID)
|
|
if err != nil {
|
|
h.respondError(c, err)
|
|
return
|
|
}
|
|
c.JSON(http.StatusOK, gin.H{"data": items})
|
|
}
|
|
|
|
// DeleteGeneration 删除当前用户的一条图片生成记录及其 COS 对象。
|
|
func (h *ProductImage) DeleteGeneration(c *gin.Context) {
|
|
taskID, err := service.ParseUUID(c.Param("task_id"), "图片生成任务")
|
|
if err != nil {
|
|
fail(c, http.StatusBadRequest, "invalid_request", err.Error())
|
|
return
|
|
}
|
|
err = h.service.DeleteGeneration(currentWebUser(c).ID, taskID, func(keys []string) error {
|
|
return h.deleteObjects(c.Request.Context(), keys)
|
|
})
|
|
if err != nil {
|
|
h.respondError(c, err)
|
|
return
|
|
}
|
|
c.Status(http.StatusNoContent)
|
|
}
|
|
|
|
// deleteObjects 删除一组已经精确解析的 COS 对象键。
|
|
func (h *ProductImage) deleteObjects(ctx context.Context, keys []string) error {
|
|
if h.cos == nil {
|
|
return nil
|
|
}
|
|
for _, key := range keys {
|
|
if err := h.cos.Delete(ctx, key); err != nil {
|
|
return err
|
|
}
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// cleanupUploadedObjects 清理提交失败前已经上传的参考图,并统一响应清理错误。
|
|
func (h *ProductImage) cleanupUploadedObjects(c *gin.Context, keys []string) bool {
|
|
if err := h.deleteObjects(c.Request.Context(), keys); err != nil {
|
|
h.respondError(c, err)
|
|
return false
|
|
}
|
|
return true
|
|
}
|
|
|
|
// cleanProductImageContentType 去除图片生成上传类型的参数部分。
|
|
func cleanProductImageContentType(value string) string {
|
|
return strings.TrimSpace(strings.Split(value, ";")[0])
|
|
}
|
|
|
|
// respondError 将图片生成模块错误转换为一致的 Web API 响应。
|
|
func (h *ProductImage) respondError(c *gin.Context, err error) {
|
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
|
fail(c, http.StatusNotFound, "not_found", "图片生成记录不存在")
|
|
return
|
|
}
|
|
message := strings.TrimSpace(err.Error())
|
|
if message == "" {
|
|
message = "请求处理失败"
|
|
}
|
|
if strings.Contains(message, "不能为空") || strings.Contains(message, "不能") || strings.Contains(message, "无效") || strings.Contains(message, "请选择") || strings.Contains(message, "最多") || strings.Contains(message, "不可用") || strings.Contains(message, "未配置") {
|
|
fail(c, http.StatusBadRequest, "invalid_request", message)
|
|
return
|
|
}
|
|
if strings.Contains(message, "积分") {
|
|
fail(c, http.StatusConflict, "insufficient_points", message)
|
|
return
|
|
}
|
|
if strings.Contains(message, "生成中") {
|
|
fail(c, http.StatusConflict, "generation_active", message)
|
|
return
|
|
}
|
|
fail(c, http.StatusInternalServerError, "internal_error", message)
|
|
}
|