// 图片生成接口处理器,负责生图参数校验、参考图上传、历史查询和结果删除。 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) }