Files
gen2d/backend/internal/service/pipeline_test.go
T
wonder 9be4b16bdc feat: 引入 Eino 框架,实现四阶段生成管线
基于 compose.Graph 编排 PromptBuilder → AssetGenerator → QualitySupervisor → FormatAdapter,
含质检重试分支(最多 3 次)和降级输出。推理层提供 mock 模式,可替换为真实 API。
2026-05-24 13:21:11 +08:00

130 lines
3.4 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package service
import (
"context"
"testing"
)
func TestPipeline_HappyPath(t *testing.T) {
// 质检一次通过
QualityChecker = func(_ context.Context, _ []GeneratedImage, _ map[string]string) (bool, string, error) {
return true, "", nil
}
defer func() { QualityChecker = defaultCheckQuality }()
output, err := RunPipeline(context.Background(), PipelineInput{
Prompt: "一个拿剑的小人",
AssetType: "sprite",
ProjectStyle: map[string]string{
"artStyle": "pixel",
"palette": "warm",
},
Params: AssetParams{
Resolution: 64,
Format: "spritesheet",
},
})
if err != nil {
t.Fatalf("RunPipeline failed: %v", err)
}
if len(output.Assets) == 0 {
t.Fatal("expected non-empty assets")
}
if output.Metadata.FrameWidth != 64 {
t.Errorf("expected FrameWidth=64, got %d", output.Metadata.FrameWidth)
}
if output.Metadata.FrameHeight != 64 {
t.Errorf("expected FrameHeight=64, got %d", output.Metadata.FrameHeight)
}
}
func TestPipeline_RetryThenPass(t *testing.T) {
// 质检前 2 次 fail,第 3 次 pass
QualityChecker = NewCountedQualityChecker(3)
defer func() { QualityChecker = defaultCheckQuality }()
output, err := RunPipeline(context.Background(), PipelineInput{
Prompt: "一把火焰剑",
AssetType: "sprite",
Params: AssetParams{
Resolution: 32,
Frames: FrameParams{
Directions: 4,
FramesPerDirection: 2,
},
Format: "spritesheet",
},
})
if err != nil {
t.Fatalf("RunPipeline failed: %v", err)
}
// 4 directions × 2 frames = 8 张图
if len(output.Assets) != 8 {
t.Errorf("expected 8 assets, got %d", len(output.Assets))
}
if output.Metadata.FrameCount != 8 {
t.Errorf("expected FrameCount=8, got %d", output.Metadata.FrameCount)
}
if output.Metadata.Directions != 4 {
t.Errorf("expected Directions=4, got %d", output.Metadata.Directions)
}
}
func TestPipeline_MaxRetryDegrade(t *testing.T) {
// 质检始终 fail,超过 3 次后降级输出
QualityChecker = AlwaysFailQualityChecker()
defer func() { QualityChecker = defaultCheckQuality }()
output, err := RunPipeline(context.Background(), PipelineInput{
Prompt: "一只飞龙",
AssetType: "sprite",
Params: AssetParams{
Resolution: 48,
},
})
if err != nil {
t.Fatalf("RunPipeline failed: %v", err)
}
// 降级也应该有输出
if len(output.Assets) == 0 {
t.Fatal("expected non-empty assets even on degrade")
}
if output.Metadata.FrameWidth != 48 {
t.Errorf("expected FrameWidth=48, got %d", output.Metadata.FrameWidth)
}
}
func TestPipeline_StyleMerge(t *testing.T) {
// 验证风格合并:task 覆盖 project
QualityChecker = func(_ context.Context, _ []GeneratedImage, style map[string]string) (bool, string, error) {
// 验证合并结果
if style["artStyle"] != "realistic" {
t.Errorf("expected artStyle=realistic (task override), got %s", style["artStyle"])
}
if style["palette"] != "warm" {
t.Errorf("expected palette=warm (from project), got %s", style["palette"])
}
return true, "", nil
}
defer func() { QualityChecker = defaultCheckQuality }()
_, err := RunPipeline(context.Background(), PipelineInput{
Prompt: "测试风格合并",
AssetType: "sprite",
ProjectStyle: map[string]string{
"artStyle": "pixel",
"palette": "warm",
},
TaskStyle: map[string]string{
"artStyle": "realistic", // 覆盖 project
},
Params: AssetParams{Resolution: 64},
})
if err != nil {
t.Fatalf("RunPipeline failed: %v", err)
}
}