Files
gen2d/backend/internal/service/nodes.go
T
Gmarker689 5d5fe60426 feat: 管线四阶段进度上报 + 图片预览 + Vite 代理
后端:
- pipeline.go: ProgressReporter 回调类型 + WithProgressReporter 注入 context
- nodes.go: 各节点 pre/post handler 调用 reportProgress() 上报阶段进度
- generate.go: TaskResponse 新增 stage 字段,runPipelineBg 注入进度回调

前端:
- vite.config.ts: 添加 /generation 代理到后端静态文件服务
- generation.ts: 轮询读取 stage 字段,暴露 stage/retryCount/rejectReason
- GeneratePage.tsx: ProgressBar 接收真实管线阶段数据
2026-05-25 14:26:26 +08:00

211 lines
6.3 KiB
Go
Executable File
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"
"strings"
"github.com/cloudwego/eino/compose"
)
// promptOptimizerNode 节点:调用 PromptAgent 生成规范提示词,合并风格与重试信息。
// 输入 PipelineInput,输出最终提示词字符串(直接供 AssetGenerator 消费)。
var promptOptimizerNode = compose.InvokableLambda(func(ctx context.Context, in PipelineInput) (string, error) {
// 合并风格描述,注入原始 Prompt 中
styleDesc := buildStyleDescription(in.ProjectStyle, in.TaskStyle)
if styleDesc != "" {
if in.Prompt != "" {
in.Prompt = in.Prompt + "。" + styleDesc
} else {
in.Prompt = styleDesc
}
}
// 注入重试原因
if in.RejectReason != "" {
if in.Prompt != "" {
in.Prompt = in.Prompt + "。注意修正以下问题:" + in.RejectReason
} else {
in.Prompt = "修正以下问题:" + in.RejectReason
}
}
if len(in.Tags) == 0 && in.Prompt == "" {
return "", fmt.Errorf("pipeline: Prompt and Tags are both 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,补上技术参数段
return appendTechNotes(in.Prompt, in.AssetType, in.Params), nil
})
// promptOptimizerPreHandler 首次运行时保存输入到 state;重试时注入 RejectReason。
func promptOptimizerPreHandler(ctx context.Context, in PipelineInput, state *PipelineState) (PipelineInput, error) {
if state.RetryCount == 0 {
state.Input = in
reportProgress(ctx, "prompt_builder", 10)
} else if state.RejectReason != "" {
in.RejectReason = state.RejectReason
reportProgress(ctx, "prompt_builder", 30+state.RetryCount*10)
}
return in, nil
}
// promptOptimizerPostHandler 将最终提示词写入全局状态。
func promptOptimizerPostHandler(ctx context.Context, out string, state *PipelineState) (string, error) {
state.FinalPrompt = out
reportProgress(ctx, "asset_generator", 35)
return out, nil
}
// AssetGenerator 节点:调用 AI 推理 API 出图。
var assetGeneratorNode = compose.InvokableLambda(func(ctx context.Context, prompt string) ([]GeneratedImage, error) {
var params AssetParams
_ = compose.ProcessState[*PipelineState](ctx, func(_ context.Context, state *PipelineState) error {
params = state.Input.Params
return nil
})
return GenerateImages(ctx, prompt, params)
})
// assetGeneratorPostHandler 将原始图片写入全局状态。
func assetGeneratorPostHandler(ctx context.Context, out []GeneratedImage, state *PipelineState) ([]GeneratedImage, error) {
state.RawImages = out
reportProgress(ctx, "quality_supervisor", 60)
return out, nil
}
// QualitySupervisor 节点:质检,设置路由目标。
var qualitySupervisorNode = compose.InvokableLambda(func(ctx context.Context, images []GeneratedImage) (PipelineInput, error) {
var input PipelineInput
err := compose.ProcessState[*PipelineState](ctx, func(_ context.Context, state *PipelineState) error {
state.RawImages = images
style := mergeStyle(state.Input.ProjectStyle, state.Input.TaskStyle)
pass, reason, checkErr := CheckQuality(ctx, images, style)
if checkErr != nil {
return fmt.Errorf("quality check: %w", checkErr)
}
state.PassQuality = pass
if !pass {
state.RejectReason = reason
}
if pass {
state.NextNode = nodeFormatAdapter
reportProgress(ctx, "format_adapter", 85)
} else if state.RetryCount >= 3 {
state.NextNode = nodeFormatAdapter
reportProgress(ctx, "format_adapter", 85)
} else {
state.RetryCount++
state.NextNode = nodePromptOptimizer
reportProgress(ctx, "quality_supervisor", 50+state.RetryCount*10)
}
input = state.Input
return nil
})
if err != nil {
return PipelineInput{}, err
}
return input, nil
})
// formatAdapterNode 节点:从 state 读取图片,格式转换,组装输出。
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 {
images = state.RawImages
return nil
})
assets := make([]Asset, len(images))
for i, img := range images {
assets[i] = Asset{
Data: img.Data,
Format: img.Format,
URL: fmt.Sprintf("output/%d.%s", i, img.Format),
}
}
resolution := input.Params.Resolution
if resolution <= 0 {
resolution = 64
}
metadata := AssetMetadata{
FrameWidth: resolution,
FrameHeight: resolution,
FrameCount: len(images),
Directions: input.Params.Frames.Directions,
}
return PipelineOutput{
Assets: assets,
Metadata: metadata,
}, nil
})
// buildStyleDescription 将风格键值对转为自然语言描述,供 PromptAgent 注入。
func buildStyleDescription(projectStyle, taskStyle map[string]string) string {
merged := mergeStyle(projectStyle, taskStyle)
if len(merged) == 0 {
return ""
}
var parts []string
for k, v := range merged {
parts = append(parts, fmt.Sprintf("%s: %s", k, v))
}
return "风格约束:" + strings.Join(parts, ";")
}
// appendTechNotes 在无标签(不走 PromptAgent)时补上技术参数段。
func appendTechNotes(prompt, assetType string, params AssetParams) string {
var parts []string
if prompt != "" {
parts = append(parts, prompt)
}
parts = append(parts, fmt.Sprintf("素材类型: %s", assetType))
if params.Resolution > 0 {
parts = append(parts, fmt.Sprintf("分辨率: %d", params.Resolution))
}
if params.Frames.Directions > 0 {
parts = append(parts, fmt.Sprintf("方向数: %d", params.Frames.Directions))
}
if params.Frames.FramesPerDirection > 0 {
parts = append(parts, fmt.Sprintf("每方向帧数: %d", params.Frames.FramesPerDirection))
}
if params.Format != "" {
parts = append(parts, fmt.Sprintf("输出格式: %s", params.Format))
}
return strings.Join(parts, ";")
}
// mergeStyle 合并工程风格与任务风格覆盖,任务同名键覆盖工程。
func mergeStyle(projectStyle, taskStyle map[string]string) map[string]string {
result := make(map[string]string)
for k, v := range projectStyle {
result[k] = v
}
for k, v := range taskStyle {
result[k] = v
}
return result
}