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 }