!67 fix: 修复管线context取消、添加FIFO任务队列、移除重复提示词优化
Merge pull request !67 from 郭永昊/fix/pipeline-queue-context-dedup
This commit is contained in:
@@ -16,6 +16,14 @@ import (
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// generateQueue 全局生成任务队列,由 main 通过 SetGenerateQueue 注入。
|
||||
var generateQueue *service.TaskQueue
|
||||
|
||||
// SetGenerateQueue 设置生成任务队列。
|
||||
func SetGenerateQueue(q *service.TaskQueue) {
|
||||
generateQueue = q
|
||||
}
|
||||
|
||||
// GenerateRequest 素材生成请求。
|
||||
type GenerateRequest struct {
|
||||
ProjectID string `json:"projectId"`
|
||||
@@ -43,7 +51,7 @@ type AssetsResponse struct {
|
||||
}
|
||||
|
||||
// Generate 素材生成接口(异步)。
|
||||
// 立即返回 taskId,后台执行管线,前端通过 GET /tasks/:taskId 轮询进度。
|
||||
// 立即返回 taskId,任务进入 FIFO 队列串行执行,前端通过 GET /tasks/:taskId 轮询进度。
|
||||
func Generate(c *gin.Context) {
|
||||
var req GenerateRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
@@ -57,18 +65,29 @@ func Generate(c *gin.Context) {
|
||||
}
|
||||
taskID := fmt.Sprintf("task-%d", time.Now().UnixMilli())
|
||||
|
||||
// 保存任务到数据库
|
||||
// 保存任务到数据库,初始状态为 pending
|
||||
if err := saveTaskToDB(c.Request.Context(), projectID, taskID, req); err != nil {
|
||||
logger.FromCtx(c.Request.Context()).Error("failed to save task", "error", err)
|
||||
c.JSON(http.StatusInternalServerError, model.Fail(http.StatusInternalServerError, "创建任务失败"))
|
||||
return
|
||||
}
|
||||
|
||||
// 返回 taskId
|
||||
c.JSON(http.StatusOK, model.OK(GenerateResponse{TaskID: taskID}))
|
||||
// 加入 FIFO 队列
|
||||
queuePos := 1
|
||||
if generateQueue != nil {
|
||||
queuePos = generateQueue.Enqueue(&service.TaskJob{
|
||||
Ctx: context.Background(),
|
||||
ProjectID: projectID,
|
||||
TaskID: taskID,
|
||||
Execute: func(ctx context.Context) error {
|
||||
runPipelineBg(ctx, projectID, taskID, req)
|
||||
return nil
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// 后台执行管线
|
||||
go runPipelineBg(c.Request.Context(), projectID, taskID, req)
|
||||
c.JSON(http.StatusOK, model.OK(GenerateResponse{TaskID: taskID}))
|
||||
_ = queuePos
|
||||
}
|
||||
|
||||
// runPipelineBg 后台执行生成管线,更新任务状态。
|
||||
@@ -82,6 +101,7 @@ func runPipelineBg(ctx context.Context, projectID, taskID string, req GenerateRe
|
||||
updateTaskInDB(ctx, taskID, "running", stage, "", progress)
|
||||
})
|
||||
|
||||
// 队列调度后才标记为 running,初始写入时是 pending
|
||||
updateTaskInDB(ctx, taskID, "running", "prompt_builder", "", 5)
|
||||
|
||||
in := service.PipelineInput{
|
||||
@@ -240,9 +260,9 @@ func saveTaskToDB(ctx context.Context, projectID, taskID string, req GenerateReq
|
||||
ProjectID: uint(projectIDUint),
|
||||
Prompt: req.Prompt,
|
||||
AssetType: req.AssetType,
|
||||
Status: "running",
|
||||
Progress: 5,
|
||||
Stage: "prompt_builder",
|
||||
Status: "pending",
|
||||
Progress: 0,
|
||||
Stage: "",
|
||||
RetryCount: 0,
|
||||
CreatedAt: time.Now(),
|
||||
UpdatedAt: time.Now(),
|
||||
|
||||
@@ -1,17 +1,23 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"fmt"
|
||||
"image/png"
|
||||
"strings"
|
||||
|
||||
"gen2d/internal/logger"
|
||||
"gen2d/pkg/gifmaker"
|
||||
"gen2d/pkg/splitsprite"
|
||||
|
||||
"github.com/cloudwego/eino/compose"
|
||||
)
|
||||
|
||||
// promptOptimizerNode 节点:调用 PromptAgent 生成规范提示词,合并风格与重试信息。
|
||||
// 输入 PipelineInput,输出最终提示词字符串(直接供 AssetGenerator 消费)。
|
||||
// promptOptimizerNode 节点:合并风格描述、重试信息与技术参数,输出最终提示词。
|
||||
// 提示词优化已由前端在提交前完成,管线内不再重复调用 PromptAgent。
|
||||
var promptOptimizerNode = compose.InvokableLambda(func(ctx context.Context, in PipelineInput) (string, error) {
|
||||
// 合并风格描述,注入原始 Prompt 中
|
||||
// 合并风格描述
|
||||
styleDesc := buildStyleDescription(in.ProjectStyle, in.TaskStyle)
|
||||
if styleDesc != "" {
|
||||
if in.Prompt != "" {
|
||||
@@ -30,26 +36,11 @@ var promptOptimizerNode = compose.InvokableLambda(func(ctx context.Context, in P
|
||||
}
|
||||
}
|
||||
|
||||
if len(in.Tags) == 0 && in.Prompt == "" {
|
||||
return "", fmt.Errorf("pipeline: Prompt and Tags are both empty")
|
||||
if in.Prompt == "" {
|
||||
return "", fmt.Errorf("pipeline: prompt is empty")
|
||||
}
|
||||
|
||||
// 有标签时调用 PromptAgent 优化提示词
|
||||
if len(in.Tags) > 0 {
|
||||
agentIn := PromptAgentInput{
|
||||
Tags: in.Tags,
|
||||
AssetType: in.AssetType,
|
||||
Prompt: in.Prompt,
|
||||
UserNote: in.UserNote,
|
||||
}
|
||||
output, err := RunPromptAgent(ctx, agentIn)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("prompt agent: %w", err)
|
||||
}
|
||||
return output.Prompt, nil
|
||||
}
|
||||
|
||||
// 无标签时直接使用原始 Prompt,补上技术参数段
|
||||
// 追加技术参数段,不再调用 PromptAgent
|
||||
return appendTechNotes(in.Prompt, in.AssetType, in.Params), nil
|
||||
})
|
||||
|
||||
@@ -127,7 +118,7 @@ var qualitySupervisorNode = compose.InvokableLambda(func(ctx context.Context, im
|
||||
return input, nil
|
||||
})
|
||||
|
||||
// formatAdapterNode 节点:从 state 读取图片,格式转换,组装输出。
|
||||
// formatAdapterNode 节点:精灵表格式时调用 splitsprite 拆分 + gifmaker 生成 GIF 预览。
|
||||
var formatAdapterNode = compose.InvokableLambda(func(ctx context.Context, input PipelineInput) (PipelineOutput, error) {
|
||||
var images []GeneratedImage
|
||||
_ = compose.ProcessState[*PipelineState](ctx, func(_ context.Context, state *PipelineState) error {
|
||||
@@ -135,6 +126,18 @@ var formatAdapterNode = compose.InvokableLambda(func(ctx context.Context, input
|
||||
return nil
|
||||
})
|
||||
|
||||
params := input.Params
|
||||
resolution := params.Resolution
|
||||
if resolution <= 0 {
|
||||
resolution = 64
|
||||
}
|
||||
|
||||
// 精灵表模式:单张图时拆分 + GIF 预览;多图时已是独立帧,透传
|
||||
if params.Format == "spritesheet" && len(images) == 1 {
|
||||
return processSpriteSheet(ctx, images[0], params, resolution)
|
||||
}
|
||||
|
||||
// 普通模式:原样透传
|
||||
assets := make([]Asset, len(images))
|
||||
for i, img := range images {
|
||||
assets[i] = Asset{
|
||||
@@ -144,23 +147,79 @@ var formatAdapterNode = compose.InvokableLambda(func(ctx context.Context, input
|
||||
}
|
||||
}
|
||||
|
||||
resolution := input.Params.Resolution
|
||||
if resolution <= 0 {
|
||||
resolution = 64
|
||||
return PipelineOutput{
|
||||
Assets: assets,
|
||||
Metadata: AssetMetadata{
|
||||
FrameWidth: resolution,
|
||||
FrameHeight: resolution,
|
||||
FrameCount: len(images),
|
||||
Directions: params.Frames.Directions,
|
||||
},
|
||||
}, nil
|
||||
})
|
||||
|
||||
// processSpriteSheet 将单张精灵表拆分为独立帧并生成 GIF 预览。
|
||||
func processSpriteSheet(ctx context.Context, img GeneratedImage, params AssetParams, resolution int) (PipelineOutput, error) {
|
||||
l := logger.FromCtx(ctx)
|
||||
src, err := png.Decode(bytes.NewReader(img.Data))
|
||||
if err != nil {
|
||||
l.Error("format_adapter decode sprite sheet failed", "error", err)
|
||||
return PipelineOutput{}, fmt.Errorf("decode sprite sheet: %w", err)
|
||||
}
|
||||
|
||||
metadata := AssetMetadata{
|
||||
FrameWidth: resolution,
|
||||
FrameHeight: resolution,
|
||||
FrameCount: len(images),
|
||||
Directions: input.Params.Frames.Directions,
|
||||
opts := splitsprite.DefaultOptions()
|
||||
if params.GridRows > 0 && params.GridCols > 0 {
|
||||
opts.GridRows = params.GridRows
|
||||
opts.GridCols = params.GridCols
|
||||
}
|
||||
|
||||
frames, err := splitsprite.Process(src, opts)
|
||||
if err != nil {
|
||||
l.Error("format_adapter split sprite sheet failed", "error", err)
|
||||
return PipelineOutput{}, fmt.Errorf("split sprite sheet: %w", err)
|
||||
}
|
||||
l.Info("format_adapter split sprite sheet", "frame_count", len(frames))
|
||||
|
||||
// 帧 → Asset
|
||||
assets := make([]Asset, 0, len(frames))
|
||||
for i, f := range frames {
|
||||
var buf bytes.Buffer
|
||||
if err := png.Encode(&buf, f); err != nil {
|
||||
l.Error("format_adapter encode frame failed", "error", err)
|
||||
return PipelineOutput{}, fmt.Errorf("encode frame %d: %w", i, err)
|
||||
}
|
||||
assets = append(assets, Asset{
|
||||
Data: buf.Bytes(),
|
||||
Format: "png",
|
||||
URL: fmt.Sprintf("output/frame_%03d.png", i),
|
||||
})
|
||||
}
|
||||
|
||||
// GIF 预览
|
||||
var gifBuf bytes.Buffer
|
||||
if err := gifmaker.Encode(&gifBuf, frames, nil); err != nil {
|
||||
l.Warn("format_adapter generate GIF preview failed", "error", err)
|
||||
} else {
|
||||
l.Info("format_adapter generated GIF preview", "size_bytes", gifBuf.Len())
|
||||
}
|
||||
|
||||
fw, fh := 0, 0
|
||||
if len(frames) > 0 {
|
||||
b := frames[0].Bounds()
|
||||
fw, fh = b.Dx(), b.Dy()
|
||||
}
|
||||
|
||||
return PipelineOutput{
|
||||
Assets: assets,
|
||||
Metadata: metadata,
|
||||
Assets: assets,
|
||||
Metadata: AssetMetadata{
|
||||
FrameWidth: fw,
|
||||
FrameHeight: fh,
|
||||
FrameCount: len(frames),
|
||||
Directions: params.Frames.Directions,
|
||||
GIFPreview: gifBuf.Bytes(),
|
||||
},
|
||||
}, nil
|
||||
})
|
||||
}
|
||||
|
||||
// buildStyleDescription 将风格键值对转为自然语言描述,供 PromptAgent 注入。
|
||||
func buildStyleDescription(projectStyle, taskStyle map[string]string) string {
|
||||
|
||||
@@ -26,7 +26,7 @@ func InitLLMConfig(cfg config.LLMConfig) {
|
||||
|
||||
// PromptAgentInput 提示词优化 Agent 输入。
|
||||
type PromptAgentInput struct {
|
||||
Tags []string `json:"tags"` // 用户选择的标签
|
||||
Tags []string `json:"tags"` // 用户选择的标签
|
||||
AssetType string `json:"assetType"` // 素材类型:sprite/background/ui/animation
|
||||
Prompt string `json:"prompt,omitempty"` // 用户原始提示词
|
||||
UserNote string `json:"userNote,omitempty"` // 用户额外描述
|
||||
@@ -89,7 +89,24 @@ func buildMetaPrompt(in PromptAgentInput) string {
|
||||
sb.WriteString("输出要求:\n")
|
||||
sb.WriteString("1. 三段式结构:【主题】描述画面主体与场景,【风格】描述艺术风格与视觉特征,【技术】描述分辨率、方向数等技术参数\n")
|
||||
sb.WriteString("2. 使用专业术语,描述具体、可执行\n")
|
||||
sb.WriteString("3. 风格一致,适合游戏资产管线\n\n")
|
||||
sb.WriteString("3. 风格一致,适合游戏资产管线\n")
|
||||
sb.WriteString("4. 根据用户标签和描述,识别素材布局模式,在【技术】段明确标注格式:\n")
|
||||
sb.WriteString(" *** 关键:帧间间隙必须留足8-16px纯白色(#FFFFFF)空白区域,间隙内不得有任何像素,确保投影法能可靠检测到间隙 ***\n")
|
||||
switch in.AssetType {
|
||||
case "sprite":
|
||||
sb.WriteString(" - 单个精灵(默认):独立PNG,纯白色背景(#FFFFFF),描述为单个角色立绘/道具图标\n")
|
||||
sb.WriteString(" - 精灵表(spritesheet):角色/道具按行列等距网格排列,纯白色背景(#FFFFFF),帧间留8-16px纯白间隙(无像素残留),标注行列数\n")
|
||||
case "background":
|
||||
sb.WriteString(" - 独立场景(默认):单张完整背景图,层次分明\n")
|
||||
sb.WriteString(" - 场景瓦片集(tileset):地形/建筑元件按规则网格排列,纯白色背景(#FFFFFF),元件间留8-16px纯白间隙,标注行列数与瓦片尺寸\n")
|
||||
case "ui":
|
||||
sb.WriteString(" - 独立UI元素(默认):单个按钮/面板/图标,纯白色背景(#FFFFFF),独立PNG\n")
|
||||
sb.WriteString(" - UI瓦片集(tileset):UI元件按规则网格排列,纯白色背景(#FFFFFF),元件间留8-16px纯白间隙,标注行列数,支持九宫格缩放\n")
|
||||
case "animation":
|
||||
sb.WriteString(" - 帧序列(默认):连续动画帧,独立帧文件或帧条带,纯白色背景(#FFFFFF)\n")
|
||||
sb.WriteString(" - 动画精灵表(spritesheet):动画帧按行列等距网格排列,纯白色背景(#FFFFFF),帧间留8-16px纯白间隙(无像素残留),标注行列数与方向数\n")
|
||||
}
|
||||
sb.WriteString("\n\n")
|
||||
|
||||
tagStr := strings.Join(in.Tags, "、")
|
||||
sb.WriteString(fmt.Sprintf("用户选择标签: %s\n", tagStr))
|
||||
@@ -242,12 +259,13 @@ func parseStreamResponse(r io.Reader) (string, error) {
|
||||
// ======================== 模板回退 ========================
|
||||
|
||||
func fallbackRefine(metaPrompt string) PromptAgentOutput {
|
||||
tags, assetType := parseTagsFromMeta(metaPrompt)
|
||||
prompt := generateStructuredPrompt(tags, assetType)
|
||||
tags, assetType, userPrompt := parseTagsFromMeta(metaPrompt)
|
||||
isSheet := isSheetRequest(tags, userPrompt)
|
||||
prompt := generateStructuredPrompt(tags, assetType, isSheet)
|
||||
return PromptAgentOutput{Prompt: prompt, RawText: prompt}
|
||||
}
|
||||
|
||||
func parseTagsFromMeta(meta string) (tags []string, assetType string) {
|
||||
func parseTagsFromMeta(meta string) (tags []string, assetType string, userPrompt string) {
|
||||
lines := strings.Split(meta, "\n")
|
||||
for _, line := range lines {
|
||||
if strings.HasPrefix(line, "用户选择标签:") {
|
||||
@@ -262,34 +280,69 @@ func parseTagsFromMeta(meta string) (tags []string, assetType string) {
|
||||
if strings.HasPrefix(line, "素材类型:") {
|
||||
assetType = strings.TrimSpace(strings.TrimPrefix(line, "素材类型: "))
|
||||
}
|
||||
if strings.HasPrefix(line, "用户原始描述:") {
|
||||
userPrompt = strings.TrimSpace(strings.TrimPrefix(line, "用户原始描述: "))
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
func generateStructuredPrompt(tags []string, assetType string) string {
|
||||
// isSheetRequest 从标签和提示词中识别是否为网格/瓦片集/精灵表模式。
|
||||
func isSheetRequest(tags []string, prompt string) bool {
|
||||
keywords := []string{"精灵表", "spritesheet", "瓦片集", "tileset", "tilemap"}
|
||||
for _, t := range tags {
|
||||
tLower := strings.ToLower(t)
|
||||
for _, kw := range keywords {
|
||||
if strings.Contains(tLower, kw) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
promptLower := strings.ToLower(prompt)
|
||||
for _, kw := range keywords {
|
||||
if strings.Contains(promptLower, kw) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func generateStructuredPrompt(tags []string, assetType string, isSheet bool) string {
|
||||
var sb strings.Builder
|
||||
sb.WriteString("【主题】")
|
||||
sb.WriteString(buildSubject(tags, assetType))
|
||||
sb.WriteString(buildSubject(tags, assetType, isSheet))
|
||||
sb.WriteString("\n")
|
||||
sb.WriteString("【风格】")
|
||||
sb.WriteString(buildStyle(tags))
|
||||
sb.WriteString("\n")
|
||||
sb.WriteString("【技术】")
|
||||
sb.WriteString(buildTechNotes(assetType))
|
||||
sb.WriteString(buildTechNotes(assetType, isSheet))
|
||||
sb.WriteString("\n")
|
||||
return sb.String()
|
||||
}
|
||||
|
||||
func buildSubject(tags []string, assetType string) string {
|
||||
func buildSubject(tags []string, assetType string, isSheet bool) string {
|
||||
tagStr := strings.Join(tags, "、")
|
||||
switch assetType {
|
||||
case "sprite":
|
||||
if isSheet {
|
||||
return fmt.Sprintf("一个融合%s元素的游戏角色精灵表,角色按行列等距网格排列,轮廓清晰,适合作为2D游戏角色", tagStr)
|
||||
}
|
||||
return fmt.Sprintf("一个融合%s元素的游戏角色精灵图,正面站立姿势,轮廓清晰,适合作为2D游戏角色", tagStr)
|
||||
case "background":
|
||||
if isSheet {
|
||||
return fmt.Sprintf("一套%s风格的游戏场景瓦片集,地形/建筑元件按规则网格排列,适合2D游戏地图拼接", tagStr)
|
||||
}
|
||||
return fmt.Sprintf("一个%s风格的游戏场景背景,层次分明,包含前景、中景和远景", tagStr)
|
||||
case "ui":
|
||||
if isSheet {
|
||||
return fmt.Sprintf("一套%s风格的UI瓦片集,UI元件按规则网格排列,适合脚本一键拆分", tagStr)
|
||||
}
|
||||
return fmt.Sprintf("一套%s风格的游戏UI元素,包括按钮、面板和图标", tagStr)
|
||||
case "animation":
|
||||
if isSheet {
|
||||
return fmt.Sprintf("一个%s风格的角色动画精灵表,动画帧按行列等距网格排列,动作流畅连贯", tagStr)
|
||||
}
|
||||
return fmt.Sprintf("一个%s风格的角色动画帧序列,动作流畅连贯", tagStr)
|
||||
default:
|
||||
return fmt.Sprintf("一个%s风格的游戏素材,高质量,适合2D游戏使用", tagStr)
|
||||
@@ -304,17 +357,29 @@ func buildStyle(tags []string) string {
|
||||
return strings.Join(parts, ";")
|
||||
}
|
||||
|
||||
func buildTechNotes(assetType string) string {
|
||||
func buildTechNotes(assetType string, isSheet bool) string {
|
||||
switch assetType {
|
||||
case "sprite":
|
||||
return "输出格式: spritesheet;分辨率: 64x64 或 128x128;透明背景"
|
||||
if isSheet {
|
||||
return "输出格式: spritesheet;帧间留8-16px纯白间隙(间隙内无任何像素),行/列间隙完全相等;纯白色背景(#FFFFFF,无渐变无噪点);标注行列数"
|
||||
}
|
||||
return "输出格式: 独立PNG;纯白色背景(#FFFFFF);尺寸: 按角色比例适配"
|
||||
case "background":
|
||||
if isSheet {
|
||||
return "输出格式: 场景瓦片集(tileset);规则网格排列;纯白色背景(#FFFFFF);元件间留8-16px纯白间隙;标注行列数与瓦片尺寸;确保无缝拼接"
|
||||
}
|
||||
return "输出格式: 独立PNG;分辨率: 1920x1080;层次分明的前中后景"
|
||||
case "ui":
|
||||
return "输出格式: 独立PNG素材;分辨率: 按元素适配;支持九宫格缩放"
|
||||
if isSheet {
|
||||
return "输出格式: UI瓦片集(tileset);规则网格排列;纯白色背景(#FFFFFF);元件间留8-16px纯白间隙;标注行列数;支持九宫格缩放;可脚本一键拆分"
|
||||
}
|
||||
return "输出格式: 独立PNG素材;纯白色背景(#FFFFFF);分辨率: 按元素适配;支持九宫格缩放"
|
||||
case "animation":
|
||||
return "输出格式: spritesheet或帧序列;建议4方向x4帧;透明背景"
|
||||
if isSheet {
|
||||
return "输出格式: 动画精灵表(spritesheet);帧间留8-16px纯白间隙(间隙内无任何像素),行/列间隙相等;纯白色背景(#FFFFFF,无渐变无噪点);标注行列数与方向数"
|
||||
}
|
||||
return "输出格式: 帧序列或帧条带;独立帧文件;纯白色背景(#FFFFFF);建议4方向x4帧"
|
||||
default:
|
||||
return "输出格式: PNG;分辨率: 标准2D游戏分辨率"
|
||||
return "输出格式: PNG;纯白色背景(#FFFFFF);分辨率: 标准2D游戏分辨率"
|
||||
}
|
||||
}
|
||||
|
||||
@@ -31,14 +31,30 @@ func TestRunPromptAgent_Fallback(t *testing.T) {
|
||||
|
||||
func TestRunPromptAgent_Sprite(t *testing.T) {
|
||||
output, err := RunPromptAgent(context.Background(), PromptAgentInput{
|
||||
Tags: []string{"像素", "中世纪", "战士"},
|
||||
Tags: []string{"像素", "中世纪", "战士", "精灵表"},
|
||||
AssetType: "sprite",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("RunPromptAgent failed: %v", err)
|
||||
}
|
||||
if !strings.Contains(output.Prompt, "spritesheet") {
|
||||
t.Errorf("sprite output should mention spritesheet: %s", output.Prompt)
|
||||
t.Errorf("sprite sheet output should mention spritesheet: %s", output.Prompt)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunPromptAgent_SpriteSingle(t *testing.T) {
|
||||
output, err := RunPromptAgent(context.Background(), PromptAgentInput{
|
||||
Tags: []string{"像素", "中世纪", "战士"},
|
||||
AssetType: "sprite",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("RunPromptAgent failed: %v", err)
|
||||
}
|
||||
if !strings.Contains(output.Prompt, "纯白色背景") {
|
||||
t.Errorf("single sprite output should mention 纯白色背景: %s", output.Prompt)
|
||||
}
|
||||
if strings.Contains(output.Prompt, "spritesheet") {
|
||||
t.Errorf("single sprite output should not mention spritesheet: %s", output.Prompt)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -251,8 +267,9 @@ func TestBuildMetaPrompt_NoUserNote(t *testing.T) {
|
||||
|
||||
func TestParseTagsFromMeta(t *testing.T) {
|
||||
meta := `用户选择标签: 像素、中世纪、战士
|
||||
素材类型: sprite`
|
||||
tags, assetType := parseTagsFromMeta(meta)
|
||||
素材类型: sprite
|
||||
用户原始描述: 一个持剑角色`
|
||||
tags, assetType, userPrompt := parseTagsFromMeta(meta)
|
||||
if len(tags) != 3 {
|
||||
t.Fatalf("expected 3 tags, got %d: %v", len(tags), tags)
|
||||
}
|
||||
@@ -262,26 +279,29 @@ func TestParseTagsFromMeta(t *testing.T) {
|
||||
if assetType != "sprite" {
|
||||
t.Errorf("expected assetType=sprite, got %s", assetType)
|
||||
}
|
||||
if userPrompt != "一个持剑角色" {
|
||||
t.Errorf("expected userPrompt='一个持剑角色', got %s", userPrompt)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseTagsFromMeta_SingleTag(t *testing.T) {
|
||||
meta := `用户选择标签: 赛博朋克
|
||||
素材类型: background`
|
||||
tags, _ := parseTagsFromMeta(meta)
|
||||
tags, _, _ := parseTagsFromMeta(meta)
|
||||
if len(tags) != 1 || tags[0] != "赛博朋克" {
|
||||
t.Errorf("expected [赛博朋克], got %v", tags)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseTagsFromMeta_Empty(t *testing.T) {
|
||||
tags, assetType := parseTagsFromMeta("no tags here")
|
||||
tags, assetType, _ := parseTagsFromMeta("no tags here")
|
||||
if len(tags) != 0 || assetType != "" {
|
||||
t.Errorf("expected empty, got tags=%v assetType=%s", tags, assetType)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerateStructuredPrompt(t *testing.T) {
|
||||
prompt := generateStructuredPrompt([]string{"像素", "战士"}, "sprite")
|
||||
prompt := generateStructuredPrompt([]string{"像素", "战士", "精灵表"}, "sprite", true)
|
||||
if !strings.HasPrefix(prompt, "【主题】") {
|
||||
t.Error("prompt should start with 【主题】")
|
||||
}
|
||||
@@ -293,21 +313,36 @@ func TestGenerateStructuredPrompt(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerateStructuredPrompt_Single(t *testing.T) {
|
||||
prompt := generateStructuredPrompt([]string{"像素", "战士"}, "sprite", false)
|
||||
if !strings.Contains(prompt, "独立PNG") {
|
||||
t.Errorf("single sprite should contain 独立PNG: %s", prompt)
|
||||
}
|
||||
if strings.Contains(prompt, "spritesheet") {
|
||||
t.Errorf("single sprite should not contain spritesheet: %s", prompt)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildSubject(t *testing.T) {
|
||||
tags := []string{"像素", "战士"}
|
||||
tests := []struct {
|
||||
assetType, want string
|
||||
isSheet bool
|
||||
}{
|
||||
{"sprite", "精灵图"},
|
||||
{"background", "场景背景"},
|
||||
{"ui", "UI元素"},
|
||||
{"animation", "动画帧序列"},
|
||||
{"unknown", "游戏素材"},
|
||||
{"sprite", "精灵图", false},
|
||||
{"sprite", "精灵表", true},
|
||||
{"background", "场景背景", false},
|
||||
{"background", "瓦片集", true},
|
||||
{"ui", "UI元素", false},
|
||||
{"ui", "瓦片集", true},
|
||||
{"animation", "动画帧序列", false},
|
||||
{"animation", "精灵表", true},
|
||||
{"unknown", "游戏素材", false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
result := buildSubject(tags, tt.assetType)
|
||||
result := buildSubject(tags, tt.assetType, tt.isSheet)
|
||||
if !strings.Contains(result, tt.want) {
|
||||
t.Errorf("buildSubject(%q) = %s, want containing %q", tt.assetType, result, tt.want)
|
||||
t.Errorf("buildSubject(%q, isSheet=%v) = %s, want containing %q", tt.assetType, tt.isSheet, result, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -327,17 +362,48 @@ func TestBuildStyle(t *testing.T) {
|
||||
func TestBuildTechNotes(t *testing.T) {
|
||||
tests := []struct {
|
||||
assetType, want string
|
||||
isSheet bool
|
||||
}{
|
||||
{"sprite", "spritesheet"},
|
||||
{"background", "1920x1080"},
|
||||
{"ui", "九宫格"},
|
||||
{"animation", "4方向x4帧"},
|
||||
{"unknown", "PNG"},
|
||||
// 默认单人模式
|
||||
{"sprite", "纯白色背景", false},
|
||||
{"background", "1920x1080", false},
|
||||
{"ui", "九宫格", false},
|
||||
{"animation", "4方向x4帧", false},
|
||||
{"unknown", "PNG", false},
|
||||
// 瓦片集/精灵表模式
|
||||
{"sprite", "spritesheet", true},
|
||||
{"background", "tileset", true},
|
||||
{"ui", "瓦片集", true},
|
||||
{"animation", "spritesheet", true},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
result := buildTechNotes(tt.assetType)
|
||||
result := buildTechNotes(tt.assetType, tt.isSheet)
|
||||
if !strings.Contains(result, tt.want) {
|
||||
t.Errorf("buildTechNotes(%q) = %s, want containing %q", tt.assetType, result, tt.want)
|
||||
t.Errorf("buildTechNotes(%q, isSheet=%v) = %s, want containing %q", tt.assetType, tt.isSheet, result, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsSheetRequest(t *testing.T) {
|
||||
tests := []struct {
|
||||
tags []string
|
||||
prompt string
|
||||
want bool
|
||||
}{
|
||||
{[]string{"像素", "精灵表"}, "", true},
|
||||
{[]string{"像素", "spritesheet"}, "", true},
|
||||
{[]string{"地形", "瓦片集"}, "", true},
|
||||
{[]string{"UI", "tileset"}, "", true},
|
||||
{[]string{"场景", "tilemap"}, "", true},
|
||||
{[]string{"像素", "战士"}, "", false},
|
||||
{[]string{"像素"}, "生成一个精灵表", true},
|
||||
{[]string{"森林"}, "场景瓦片集", true},
|
||||
{nil, "", false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
got := isSheetRequest(tt.tags, tt.prompt)
|
||||
if got != tt.want {
|
||||
t.Errorf("isSheetRequest(tags=%v, prompt=%q) = %v, want %v", tt.tags, tt.prompt, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,104 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
|
||||
"gen2d/internal/logger"
|
||||
)
|
||||
|
||||
// TaskJob 队列中的任务。
|
||||
type TaskJob struct {
|
||||
Ctx context.Context
|
||||
ProjectID string
|
||||
TaskID string
|
||||
Execute func(ctx context.Context) error
|
||||
}
|
||||
|
||||
// TaskQueue 串行 FIFO 任务队列。
|
||||
type TaskQueue struct {
|
||||
mu sync.Mutex
|
||||
jobs []*TaskJob
|
||||
ready chan struct{}
|
||||
stop chan struct{}
|
||||
stopped bool
|
||||
}
|
||||
|
||||
// NewTaskQueue 创建任务队列并启动调度协程。
|
||||
func NewTaskQueue() *TaskQueue {
|
||||
q := &TaskQueue{
|
||||
jobs: make([]*TaskJob, 0),
|
||||
ready: make(chan struct{}, 1),
|
||||
stop: make(chan struct{}),
|
||||
}
|
||||
go q.run()
|
||||
return q
|
||||
}
|
||||
|
||||
// Enqueue 将任务加入队尾,返回队列中的位置(1-based)。
|
||||
func (q *TaskQueue) Enqueue(job *TaskJob) int {
|
||||
q.mu.Lock()
|
||||
defer q.mu.Unlock()
|
||||
q.jobs = append(q.jobs, job)
|
||||
pos := len(q.jobs)
|
||||
select {
|
||||
case q.ready <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
return pos
|
||||
}
|
||||
|
||||
// QueueLen 返回当前队列长度。
|
||||
func (q *TaskQueue) QueueLen() int {
|
||||
q.mu.Lock()
|
||||
defer q.mu.Unlock()
|
||||
return len(q.jobs)
|
||||
}
|
||||
|
||||
// Stop 优雅关闭队列。
|
||||
func (q *TaskQueue) Stop() {
|
||||
q.mu.Lock()
|
||||
defer q.mu.Unlock()
|
||||
if !q.stopped {
|
||||
q.stopped = true
|
||||
close(q.stop)
|
||||
}
|
||||
}
|
||||
|
||||
func (q *TaskQueue) run() {
|
||||
for {
|
||||
select {
|
||||
case <-q.stop:
|
||||
return
|
||||
case <-q.ready:
|
||||
q.processNext()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (q *TaskQueue) processNext() {
|
||||
q.mu.Lock()
|
||||
if len(q.jobs) == 0 {
|
||||
q.mu.Unlock()
|
||||
return
|
||||
}
|
||||
job := q.jobs[0]
|
||||
q.jobs = q.jobs[1:]
|
||||
// 如果队列还有任务,重新发信号
|
||||
if len(q.jobs) > 0 {
|
||||
select {
|
||||
case q.ready <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
q.mu.Unlock()
|
||||
|
||||
l := logger.With("task_id", job.TaskID, "project_id", job.ProjectID)
|
||||
l.Info("task queue executing job", "queue_remaining", len(q.jobs))
|
||||
|
||||
if err := job.Execute(job.Ctx); err != nil {
|
||||
l.Error("task job failed", "error", err)
|
||||
} else {
|
||||
l.Info("task job completed")
|
||||
}
|
||||
}
|
||||
@@ -36,11 +36,14 @@ type AssetParams struct {
|
||||
Resolution int
|
||||
Frames FrameParams
|
||||
Format string // "spritesheet" / "individual"
|
||||
// GridRows / GridCols override projection-based split for sprite sheets.
|
||||
GridRows int
|
||||
GridCols int
|
||||
}
|
||||
|
||||
// FrameParams 帧参数
|
||||
type FrameParams struct {
|
||||
Directions int
|
||||
Directions int
|
||||
FramesPerDirection int
|
||||
}
|
||||
|
||||
@@ -65,4 +68,5 @@ type AssetMetadata struct {
|
||||
FrameHeight int
|
||||
FrameCount int
|
||||
Directions int
|
||||
GIFPreview []byte `json:"-"` // animated GIF preview (not serialized)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user