Files
JuYou/API/internal/handler/product_image.go
T
2026-08-25 17:59:42 +08:00

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)
}