Files
gen2d/backend/internal/service/pipeline.go
T
Gmarker689 9ca1e65221 feat(pipeline): PromptOptimizer 集成到生成管线,新增优化 API
- PipelineInput 新增 Tags/UserNote 字段
- 管线重构为: START → PromptOptimizer → AssetGenerator → QualitySupervisor → FormatAdapter
- promptOptimizerNode: 有标签调用 PromptAgent,无标签补技术参数,合并风格与重试信息
- PromptBuilder 移除,提示词构建逻辑并入 PromptOptimizer
- 新增 POST /api/v1/prompt/optimize 接口
- main.go 注入 LLM/ImageGen 配置,注册 prompt 路由
- inference.go 新增 InitImageGenConfig 注入
2026-05-24 20:45:18 +08:00

102 lines
3.5 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"
"fmt"
"github.com/cloudwego/eino/compose"
)
const (
nodePromptOptimizer = "prompt_optimizer"
nodeAssetGenerator = "asset_generator"
nodeQualitySupervisor = "quality_supervisor"
nodeFormatAdapter = "format_adapter"
)
// NewGenerateGraph 创建生成管线 Graph(PromptOptimizer → AssetGenerator → QualitySupervisor → FormatAdapter)。
//
// START → PromptOptimizer → AssetGenerator → QualitySupervisor
// ├── pass → FormatAdapter → END
// └── fail, retry<3 → PromptOptimizer
// └── fail, retry>=3 → FormatAdapter (降级)
func NewGenerateGraph() (*compose.Graph[PipelineInput, PipelineOutput], error) {
g := compose.NewGraph[PipelineInput, PipelineOutput](
compose.WithGenLocalState(func(ctx context.Context) *PipelineState {
return &PipelineState{}
}),
)
if err := g.AddLambdaNode(nodePromptOptimizer, promptOptimizerNode,
compose.WithStatePreHandler(promptOptimizerPreHandler),
compose.WithStatePostHandler(promptOptimizerPostHandler),
); err != nil {
return nil, fmt.Errorf("add %s node: %w", nodePromptOptimizer, err)
}
if err := g.AddLambdaNode(nodeAssetGenerator, assetGeneratorNode,
compose.WithStatePostHandler(assetGeneratorPostHandler),
); err != nil {
return nil, fmt.Errorf("add %s node: %w", nodeAssetGenerator, err)
}
if err := g.AddLambdaNode(nodeQualitySupervisor, qualitySupervisorNode); err != nil {
return nil, fmt.Errorf("add %s node: %w", nodeQualitySupervisor, err)
}
if err := g.AddLambdaNode(nodeFormatAdapter, formatAdapterNode); err != nil {
return nil, fmt.Errorf("add %s node: %w", nodeFormatAdapter, err)
}
// 连线:START → PromptOptimizer → AssetGenerator → Supervisor
if err := g.AddEdge(compose.START, nodePromptOptimizer); err != nil {
return nil, fmt.Errorf("add edge START->%s: %w", nodePromptOptimizer, err)
}
if err := g.AddEdge(nodePromptOptimizer, nodeAssetGenerator); err != nil {
return nil, fmt.Errorf("add edge %s->%s: %w", nodePromptOptimizer, nodeAssetGenerator, err)
}
if err := g.AddEdge(nodeAssetGenerator, nodeQualitySupervisor); err != nil {
return nil, fmt.Errorf("add edge %s->%s: %w", nodeAssetGenerator, nodeQualitySupervisor, err)
}
if err := g.AddEdge(nodeFormatAdapter, compose.END); err != nil {
return nil, fmt.Errorf("add edge %s->END: %w", nodeFormatAdapter, err)
}
// 连线:质检分支(从 state.NextNode 读取路由目标)
if err := g.AddBranch(nodeQualitySupervisor, compose.NewGraphBranch(
func(ctx context.Context, _ PipelineInput) (string, error) {
var next string
_ = compose.ProcessState[*PipelineState](ctx, func(_ context.Context, state *PipelineState) error {
next = state.NextNode
return nil
})
return next, nil
},
map[string]bool{nodePromptOptimizer: true, nodeFormatAdapter: true},
)); err != nil {
return nil, fmt.Errorf("add branch at %s: %w", nodeQualitySupervisor, err)
}
return g, nil
}
// RunPipeline 编译并执行生成管线。
func RunPipeline(ctx context.Context, in PipelineInput) (*PipelineOutput, error) {
g, err := NewGenerateGraph()
if err != nil {
return nil, fmt.Errorf("create graph: %w", err)
}
r, err := g.Compile(ctx, compose.WithMaxRunSteps(20))
if err != nil {
return nil, fmt.Errorf("compile graph: %w", err)
}
output, err := r.Invoke(ctx, in)
if err != nil {
return nil, fmt.Errorf("invoke pipeline: %w", err)
}
return &output, nil
}