初始化

This commit is contained in:
Ran
2026-08-25 17:59:42 +08:00
commit 4b7380dd9b
408 changed files with 327400 additions and 0 deletions
@@ -0,0 +1,323 @@
// 图片生成生图业务模块,负责提交用户级图片任务、查询历史并清理生成资源。
package productimage
import (
"encoding/json"
"errors"
"fmt"
"strings"
"time"
"unicode/utf8"
"juhe-factory/api/internal/billing"
"juhe-factory/api/internal/model"
queuepkg "juhe-factory/api/internal/queue"
"github.com/google/uuid"
"github.com/hibiken/asynq"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
var activeTaskStatuses = []string{"pending_submission", "submitting", "submitted", "processing", "result_ready", "downloading", "cancel_requested"}
// ReferenceUpload 表示已经上传到对象存储、等待绑定到生图任务的参考图。
type ReferenceUpload struct {
Asset *model.MediaAsset
Name string
}
// GenerateInput 是图片生成生图接口的配置参数和临时参考图。
type GenerateInput struct {
Prompt string
ModelID uuid.UUID
AspectRatio string
Resolution string
References []ReferenceUpload
}
// GenerationView 汇总任务状态和已经落库的图片结果。
type GenerationView struct {
TaskID uuid.UUID `json:"task_id"`
Status string `json:"status"`
Prompt string `json:"prompt"`
ModelName string `json:"model_name"`
ErrorMessage string `json:"error_message"`
MediaAssetID *uuid.UUID `json:"media_asset_id"`
PublicURL string `json:"public_url"`
MimeType string `json:"mime_type"`
EstimatedPoints string `json:"estimated_points"`
ActualPoints *string `json:"actual_points"`
CreatedAt time.Time `json:"created_at"`
FinishedAt *time.Time `json:"finished_at"`
}
// Service 组合图片生成模块需要的数据库和任务队列依赖。
type Service struct {
DB *gorm.DB
Queue *asynq.Client
}
// NewService 创建图片生成生图服务。
func NewService(db *gorm.DB, queue *asynq.Client) *Service { return &Service{DB: db, Queue: queue} }
// QueueGeneration 校验配置、保存临时参考图并提交一个独立图片任务。
func (s *Service) QueueGeneration(userID uuid.UUID, input GenerateInput) (*model.GenerationTask, error) {
if s.Queue == nil {
return nil, errors.New("生成任务队列不可用")
}
prompt := strings.TrimSpace(input.Prompt)
if prompt == "" {
return nil, errors.New("提示词不能为空")
}
if utf8.RuneCountInString(prompt) > 5000 {
return nil, errors.New("提示词不能超过 5000 个字符")
}
if len(input.References) > 4 {
return nil, errors.New("最多上传 4 张参考图")
}
if input.ModelID == uuid.Nil {
return nil, errors.New("请选择图片模型")
}
if !validRatio(input.AspectRatio) || !validResolution(input.Resolution) {
return nil, errors.New("图片生成参数无效")
}
var selected struct {
ModelID uuid.UUID
ChannelID uuid.UUID
ModelName string
Price string
PriceExists bool
}
if err := s.DB.Raw(`SELECT model.id AS model_id,model.channel_id,model.name AS model_name,
coalesce(price.price,0)::text AS price,(price.id IS NOT NULL) AS price_exists
FROM models model JOIN channels channel ON channel.id=model.channel_id
LEFT JOIN model_prices price ON price.model_id=model.id AND lower(price.price_key)=lower(?)
WHERE model.id=? AND model.model_type='image' AND model.enabled=true AND model.deleted_at IS NULL
AND channel.enabled=true AND channel.deleted_at IS NULL`, input.Resolution, input.ModelID).Scan(&selected).Error; err != nil {
return nil, err
}
if selected.ModelID == uuid.Nil {
return nil, errors.New("图片模型不可用")
}
if !selected.PriceExists {
return nil, fmt.Errorf("当前图片模型未配置 %s 价格", input.Resolution)
}
imageURLs := make([]string, 0, len(input.References))
referenceIDs := make([]string, 0, len(input.References))
referenceNames := make([]string, 0, len(input.References))
seenNames := make(map[string]struct{}, len(input.References))
for _, reference := range input.References {
if reference.Asset == nil || strings.TrimSpace(reference.Asset.PublicURL) == "" {
return nil, errors.New("参考图信息无效")
}
name := strings.TrimSpace(reference.Name)
if name == "" {
return nil, errors.New("参考图名称不能为空")
}
if _, exists := seenNames[name]; exists {
return nil, errors.New("参考图名称不能重复")
}
seenNames[name] = struct{}{}
imageURLs = append(imageURLs, reference.Asset.PublicURL)
referenceIDs = append(referenceIDs, reference.Asset.ID.String())
referenceNames = append(referenceNames, name)
}
providerPrompt := buildProviderPrompt(prompt, referenceNames)
payload, err := json.Marshal(map[string]any{
"product_image": true, "prompt": providerPrompt, "display_prompt": prompt,
"model": selected.ModelName, "size": input.AspectRatio, "resolution": input.Resolution,
"image_urls": imageURLs, "reference_media_asset_ids": referenceIDs, "reference_names": referenceNames, "n": 1,
})
if err != nil {
return nil, err
}
task := &model.GenerationTask{ID: uuid.New(), RequestID: "product_image_" + uuid.NewString(), UserID: userID, ChannelID: &selected.ChannelID, ModelID: &selected.ModelID, TaskType: "image_generation", Status: "pending_submission", InputData: payload}
err = s.DB.Transaction(func(tx *gorm.DB) error {
for _, reference := range input.References {
owner := userID
reference.Asset.OwnerUserID = &owner
if err := tx.Create(reference.Asset).Error; err != nil {
return err
}
}
return billing.PrechargeGenerationTask(tx, task, selected.Price, 1, "图片生成")
})
if err != nil {
return nil, err
}
if err := queuepkg.EnqueueID(s.Queue, queuepkg.TypeDispatchChannel, selected.ChannelID, 0); err != nil {
if refundErr := s.failQueuedTask(task.ID, err.Error()); refundErr != nil {
return nil, fmt.Errorf("生成任务入队失败且预扣返还失败: %w", refundErr)
}
return nil, errors.New("生成任务队列暂时不可用,请稍后重试")
}
return task, nil
}
// Generations 返回当前用户最近五十条图片生成记录。
func (s *Service) Generations(userID uuid.UUID) ([]GenerationView, error) {
items := make([]GenerationView, 0)
err := s.DB.Raw(`SELECT task.id AS task_id,task.status,
coalesce(task.input_data->>'display_prompt',task.input_data->>'prompt','') AS prompt,
coalesce(model.name,task.input_data->>'model','') AS model_name,task.error_message,
output.media_asset_id,coalesce(media.public_url,'') AS public_url,coalesce(media.mime_type,'') AS mime_type,
task.estimated_points,task.actual_points,task.created_at,task.finished_at
FROM generation_tasks task
LEFT JOIN models model ON model.id=task.model_id
LEFT JOIN generation_outputs output ON output.task_id=task.id AND output.sequence_no=1
LEFT JOIN media_assets media ON media.id=output.media_asset_id AND media.deleted_at IS NULL
WHERE task.user_id=? AND task.task_type='image_generation' AND task.input_data->>'product_image'='true'
ORDER BY task.created_at DESC LIMIT 50`, userID).Scan(&items).Error
return items, err
}
// DeleteGeneration 清理当前用户的图片任务、生成媒体和参考媒体;活动任务也允许主动删除。
func (s *Service) DeleteGeneration(userID, taskID uuid.UUID, deleteObjects func([]string) error) error {
objectKeys := make([]string, 0)
var channelID *uuid.UUID
err := s.DB.Transaction(func(tx *gorm.DB) error {
var task model.GenerationTask
if err := tx.Where("id=? AND user_id=? AND task_type='image_generation' AND input_data->>'product_image'='true'", taskID, userID).Take(&task).Error; err != nil {
return err
}
if !canDeleteGeneration(task.Status, task.ErrorMessage) {
return errors.New("图片仍在正常生成中,暂时不能删除")
}
channelID = task.ChannelID
var input struct {
ReferenceMediaAssetIDs []uuid.UUID `json:"reference_media_asset_ids"`
}
_ = json.Unmarshal(task.InputData, &input)
mediaIDs := append([]uuid.UUID(nil), input.ReferenceMediaAssetIDs...)
var outputs []struct {
MediaAssetID uuid.UUID
ObjectKey string
}
if err := tx.Table("generation_outputs output").Select("output.media_asset_id,media.object_key").Joins("JOIN media_assets media ON media.id=output.media_asset_id").Where("output.task_id=?", taskID).Scan(&outputs).Error; err != nil {
return err
}
for _, output := range outputs {
mediaIDs = append(mediaIDs, output.MediaAssetID)
if output.ObjectKey != "" {
objectKeys = append(objectKeys, output.ObjectKey)
}
}
if len(mediaIDs) > 0 {
var refs []struct{ ObjectKey string }
if err := tx.Table("media_assets").Select("object_key").Where("id IN ?", mediaIDs).Scan(&refs).Error; err != nil {
return err
}
for _, ref := range refs {
if ref.ObjectKey != "" {
objectKeys = append(objectKeys, ref.ObjectKey)
}
}
}
if deleteObjects != nil {
if err := deleteObjects(uniqueStrings(objectKeys)); err != nil {
return err
}
}
if err := tx.Exec("DELETE FROM generation_outputs WHERE task_id=?", taskID).Error; err != nil {
return err
}
if len(mediaIDs) > 0 {
if err := tx.Exec("DELETE FROM channel_asset_cache WHERE media_asset_id IN ?", mediaIDs).Error; err != nil {
return err
}
if err := tx.Exec("DELETE FROM media_assets WHERE id IN ?", mediaIDs).Error; err != nil {
return err
}
}
return tx.Exec("DELETE FROM generation_tasks WHERE id=?", taskID).Error
})
if err == nil && channelID != nil && s.Queue != nil {
_ = queuepkg.EnqueueID(s.Queue, queuepkg.TypeDispatchChannel, *channelID, 0)
}
return err
}
// canDeleteGeneration 判断图片任务是否已结束,或因错误停留在活动状态而允许用户清理。
func canDeleteGeneration(status, errorMessage string) bool {
for _, activeStatus := range activeTaskStatuses {
if status == activeStatus {
return strings.TrimSpace(errorMessage) != ""
}
}
return true
}
// failQueuedTask 在调度队列不可用时结束任务、返还积分并删除参考媒体记录。
func (s *Service) failQueuedTask(taskID uuid.UUID, message string) error {
return s.DB.Transaction(func(tx *gorm.DB) error {
var task model.GenerationTask
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Where("id=?", taskID).Take(&task).Error; err != nil {
return err
}
refunded, err := billing.RefundGenerationTask(tx, &task, "图片生成入队失败返还")
if err != nil {
return err
}
var input struct {
ReferenceMediaAssetIDs []uuid.UUID `json:"reference_media_asset_ids"`
}
_ = json.Unmarshal(task.InputData, &input)
if len(input.ReferenceMediaAssetIDs) > 0 {
if err := tx.Exec("DELETE FROM media_assets WHERE id IN ?", input.ReferenceMediaAssetIDs).Error; err != nil {
return err
}
}
updates := map[string]any{"status": "failed", "error_code": "queue_unavailable", "error_message": message, "finished_at": time.Now(), "actual_points": "0.00"}
if refunded {
updates["cost_refunded"] = true
}
return tx.Model(&model.GenerationTask{}).Where("id=?", taskID).Updates(updates).Error
})
}
// validRatio 判断画幅比例是否属于图片生成工具公开支持的取值。
func validRatio(value string) bool {
for _, allowed := range []string{"9:21", "1:1", "3:4", "4:3", "9:16", "16:9"} {
if value == allowed {
return true
}
}
return false
}
// validResolution 判断分辨率是否属于图片生成工具公开支持的档位。
func validResolution(value string) bool {
for _, allowed := range []string{"1k", "2k", "4k"} {
if value == allowed {
return true
}
}
return false
}
// buildProviderPrompt 将文件名提及与参考图数组顺序建立明确对应关系。
func buildProviderPrompt(prompt string, names []string) string {
if len(names) == 0 {
return prompt
}
mappings := make([]string, 0, len(names))
for index, name := range names {
mappings = append(mappings, fmt.Sprintf("第%d张参考图名称为“%s”", index+1, name))
}
return "参考图对应关系:" + strings.Join(mappings, "") + "。\n" + prompt
}
// uniqueStrings 按首次出现顺序去除空对象键和重复对象键。
func uniqueStrings(values []string) []string {
seen := make(map[string]struct{}, len(values))
result := make([]string, 0, len(values))
for _, value := range values {
if _, ok := seen[value]; ok || value == "" {
continue
}
seen[value] = struct{}{}
result = append(result, value)
}
return result
}
@@ -0,0 +1,56 @@
// 图片生成服务测试,验证参考图文件名映射和对象键去重规则。
package productimage
import "testing"
// TestBuildProviderPrompt 验证文件名提及会与参考图数组顺序建立稳定对应关系。
func TestBuildProviderPrompt(t *testing.T) {
t.Parallel()
prompt := "让@正面图.png中的商品使用@包装图.webp的包装"
got := buildProviderPrompt(prompt, []string{"正面图.png", "包装图.webp"})
want := "参考图对应关系:第1张参考图名称为“正面图.png”;第2张参考图名称为“包装图.webp”。\n" + prompt
if got != want {
t.Fatalf("buildProviderPrompt() = %q, want %q", got, want)
}
if withoutReferences := buildProviderPrompt(prompt, nil); withoutReferences != prompt {
t.Fatalf("buildProviderPrompt() without references = %q, want %q", withoutReferences, prompt)
}
}
// TestUniqueStrings 验证重复或空对象键不会触发重复删除请求。
func TestUniqueStrings(t *testing.T) {
t.Parallel()
got := uniqueStrings([]string{"a", "", "b", "a", "b", "c"})
want := []string{"a", "b", "c"}
if len(got) != len(want) {
t.Fatalf("uniqueStrings() length = %d, want %d", len(got), len(want))
}
for index := range want {
if got[index] != want[index] {
t.Fatalf("uniqueStrings()[%d] = %q, want %q", index, got[index], want[index])
}
}
}
// TestCanDeleteGeneration 验证失败停滞的活动任务可删除,而正常生成中的任务仍受保护。
func TestCanDeleteGeneration(t *testing.T) {
t.Parallel()
tests := []struct {
name string
status string
errorMessage string
want bool
}{
{name: "正常生成", status: "processing", want: false},
{name: "连接超时", status: "pending_submission", errorMessage: "dial tcp: i/o timeout", want: true},
{name: "生成失败", status: "failed", want: true},
{name: "生成成功", status: "succeeded", want: true},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
if got := canDeleteGeneration(test.status, test.errorMessage); got != test.want {
t.Fatalf("canDeleteGeneration(%q, %q) = %t, want %t", test.status, test.errorMessage, got, test.want)
}
})
}
}