From 788cd021efe97655de5a9c1066accd6fee2ba648 Mon Sep 17 00:00:00 2001 From: Gmaker689 <1711322114@qq.com> Date: Sun, 24 May 2026 20:45:09 +0800 Subject: [PATCH] =?UTF-8?q?feat(prompt-agent):=20=E6=96=B0=E5=A2=9E=20Prom?= =?UTF-8?q?ptAgent=20Eino=20Chain=20=E8=B0=83=E7=94=A8=20LLM=20=E4=BC=98?= =?UTF-8?q?=E5=8C=96=E6=8F=90=E7=A4=BA=E8=AF=8D?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Eino Chain: formatMetaPrompt → llmRefine → 三段式输出 - callLLMRefine 调用 OpenAI 兼容 Chat Completions API - 支持 SSE 流式与非流式响应解析 - 无 API key 或调用失败时自动回退模板生成 - 新增完整的单元测试与集成测试(含 mock HTTP server) --- backend/internal/service/prompt_agent.go | 313 ++++++++++++++++ backend/internal/service/prompt_agent_test.go | 343 ++++++++++++++++++ 2 files changed, 656 insertions(+) create mode 100644 backend/internal/service/prompt_agent.go create mode 100644 backend/internal/service/prompt_agent_test.go diff --git a/backend/internal/service/prompt_agent.go b/backend/internal/service/prompt_agent.go new file mode 100644 index 0000000..813ade2 --- /dev/null +++ b/backend/internal/service/prompt_agent.go @@ -0,0 +1,313 @@ +package service + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "log" + "net/http" + "strings" + + "gen2d/internal/config" + + "github.com/cloudwego/eino/compose" +) + +// llmCfg 保存 LLM 配置,由 main 通过 InitLLMConfig 注入。 +var llmCfg config.LLMConfig + +// InitLLMConfig 注入 LLM 配置。 +func InitLLMConfig(cfg config.LLMConfig) { + llmCfg = cfg +} + +// PromptAgentInput 提示词优化 Agent 输入。 +type PromptAgentInput struct { + Tags []string `json:"tags"` // 用户选择的标签 + AssetType string `json:"assetType"` // 素材类型:sprite/background/ui/animation + Prompt string `json:"prompt,omitempty"` // 用户原始提示词 + UserNote string `json:"userNote,omitempty"` // 用户额外描述 +} + +// PromptAgentOutput 提示词优化 Agent 输出。 +type PromptAgentOutput struct { + Prompt string `json:"prompt"` // 优化后的标准提示词 + RawText string `json:"rawText"` // LLM 原始返回文本 +} + +// NewPromptAgentGraph 创建提示词优化 Agent Chain。 +// +// START → formatMetaPrompt → llmRefine → END +func NewPromptAgentGraph() (*compose.Chain[PromptAgentInput, PromptAgentOutput], error) { + chain := compose.NewChain[PromptAgentInput, PromptAgentOutput]() + + chain.AppendLambda( + compose.InvokableLambda(func(ctx context.Context, in PromptAgentInput) (string, error) { + return buildMetaPrompt(in), nil + }), + ) + + chain.AppendLambda( + compose.InvokableLambda(func(ctx context.Context, metaPrompt string) (PromptAgentOutput, error) { + return callLLMRefine(ctx, metaPrompt) + }), + ) + + return chain, nil +} + +// RunPromptAgent 编译并执行提示词优化 Agent。 +func RunPromptAgent(ctx context.Context, in PromptAgentInput) (*PromptAgentOutput, error) { + g, err := NewPromptAgentGraph() + if err != nil { + return nil, fmt.Errorf("create prompt agent: %w", err) + } + + r, err := g.Compile(ctx) + if err != nil { + return nil, fmt.Errorf("compile prompt agent: %w", err) + } + + output, err := r.Invoke(ctx, in) + if err != nil { + return nil, fmt.Errorf("invoke prompt agent: %w", err) + } + + return &output, nil +} + +// buildMetaPrompt 构建发送给 LLM 的元提示词。 +func buildMetaPrompt(in PromptAgentInput) string { + var sb strings.Builder + + sb.WriteString("你是一个专业的 2D 游戏素材提示词工程师。\n") + sb.WriteString("根据用户提供的标签和素材类型,生成一个详细、规范的中文提示词。\n\n") + + sb.WriteString("输出要求:\n") + sb.WriteString("1. 三段式结构:【主题】描述画面主体与场景,【风格】描述艺术风格与视觉特征,【技术】描述分辨率、方向数等技术参数\n") + sb.WriteString("2. 使用专业术语,描述具体、可执行\n") + sb.WriteString("3. 风格一致,适合游戏资产管线\n\n") + + tagStr := strings.Join(in.Tags, "、") + sb.WriteString(fmt.Sprintf("用户选择标签: %s\n", tagStr)) + sb.WriteString(fmt.Sprintf("素材类型: %s\n", in.AssetType)) + if in.Prompt != "" { + sb.WriteString(fmt.Sprintf("用户原始描述: %s\n", in.Prompt)) + } + if in.UserNote != "" { + sb.WriteString(fmt.Sprintf("补充说明: %s\n", in.UserNote)) + } + + return sb.String() +} + +// ======================== ChatModel 调用层 ======================== + +// chatMessage OpenAI 兼容的消息结构。 +type chatMessage struct { + Role string `json:"role"` + Content string `json:"content"` +} + +// chatRequest OpenAI 兼容的请求体。 +type chatRequest struct { + Model string `json:"model"` + Messages []chatMessage `json:"messages"` + Temperature float64 `json:"temperature"` + MaxTokens int `json:"max_tokens,omitempty"` + Stream bool `json:"stream"` +} + +// chatResponse OpenAI 兼容的非流式响应。 +type chatResponse struct { + Choices []struct { + Message chatMessage `json:"message"` + } `json:"choices"` +} + +// callLLMRefine 调用 LLM 生成规范化提示词。未配置 API key 时回退到模板生成。 +func callLLMRefine(ctx context.Context, metaPrompt string) (PromptAgentOutput, error) { + if llmCfg.APIKey == "" { + log.Println("[prompt_agent] LLM API key not configured, using template fallback") + return fallbackRefine(metaPrompt), nil + } + + messages := []chatMessage{ + {Role: "system", Content: "你是一个专业的2D游戏素材提示词工程师,只输出优化后的中文提示词,不要任何解释。"}, + {Role: "user", Content: metaPrompt}, + } + + result, err := chatCompletion(ctx, messages) + if err != nil { + log.Printf("[prompt_agent] LLM API call failed: %v, using template fallback", err) + return fallbackRefine(metaPrompt), nil + } + + return PromptAgentOutput{ + Prompt: result, + RawText: result, + }, nil +} + +// chatCompletion 调用 OpenAI 兼容的 Chat Completions API。 +func chatCompletion(ctx context.Context, messages []chatMessage) (string, error) { + reqBody := chatRequest{ + Model: llmCfg.Model, + Messages: messages, + Temperature: llmCfg.Temperature, + MaxTokens: llmCfg.MaxTokens, + Stream: false, + } + + body, err := json.Marshal(reqBody) + if err != nil { + return "", fmt.Errorf("marshal request: %w", err) + } + + url := strings.TrimRight(llmCfg.BaseURL, "/") + "/chat/completions" + req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body)) + if err != nil { + return "", fmt.Errorf("create request: %w", err) + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer "+llmCfg.APIKey) + + resp, err := http.DefaultClient.Do(req) + if err != nil { + return "", fmt.Errorf("send request: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + b, _ := io.ReadAll(io.LimitReader(resp.Body, 4096)) + return "", fmt.Errorf("llm api error %d: %s", resp.StatusCode, string(b)) + } + + // 处理流式 (SSE) 与非流式响应 + if resp.Header.Get("Content-Type") == "text/event-stream" { + return parseStreamResponse(resp.Body) + } + + var chatResp chatResponse + if err := json.NewDecoder(resp.Body).Decode(&chatResp); err != nil { + return "", fmt.Errorf("decode response: %w", err) + } + if len(chatResp.Choices) == 0 { + return "", fmt.Errorf("llm returned empty choices") + } + + return chatResp.Choices[0].Message.Content, nil +} + +// parseStreamResponse 解析 SSE 流式响应。 +func parseStreamResponse(r io.Reader) (string, error) { + var sb strings.Builder + scanner := bufio.NewScanner(r) + for scanner.Scan() { + line := scanner.Text() + if !strings.HasPrefix(line, "data: ") { + continue + } + data := strings.TrimPrefix(line, "data: ") + if data == "[DONE]" { + break + } + var chunk struct { + Choices []struct { + Delta struct { + Content string `json:"content"` + } `json:"delta"` + } `json:"choices"` + } + if err := json.Unmarshal([]byte(data), &chunk); err != nil { + continue + } + for _, c := range chunk.Choices { + sb.WriteString(c.Delta.Content) + } + } + return sb.String(), scanner.Err() +} + +// ======================== 模板回退 ======================== + +func fallbackRefine(metaPrompt string) PromptAgentOutput { + tags, assetType := parseTagsFromMeta(metaPrompt) + prompt := generateStructuredPrompt(tags, assetType) + return PromptAgentOutput{Prompt: prompt, RawText: prompt} +} + +func parseTagsFromMeta(meta string) (tags []string, assetType string) { + lines := strings.Split(meta, "\n") + for _, line := range lines { + if strings.HasPrefix(line, "用户选择标签:") { + tagStr := strings.TrimPrefix(line, "用户选择标签: ") + for _, t := range strings.Split(tagStr, "、") { + t = strings.TrimSpace(t) + if t != "" { + tags = append(tags, t) + } + } + } + if strings.HasPrefix(line, "素材类型:") { + assetType = strings.TrimSpace(strings.TrimPrefix(line, "素材类型: ")) + } + } + return +} + +func generateStructuredPrompt(tags []string, assetType string) string { + var sb strings.Builder + sb.WriteString("【主题】") + sb.WriteString(buildSubject(tags, assetType)) + sb.WriteString("\n") + sb.WriteString("【风格】") + sb.WriteString(buildStyle(tags)) + sb.WriteString("\n") + sb.WriteString("【技术】") + sb.WriteString(buildTechNotes(assetType)) + sb.WriteString("\n") + return sb.String() +} + +func buildSubject(tags []string, assetType string) string { + tagStr := strings.Join(tags, "、") + switch assetType { + case "sprite": + return fmt.Sprintf("一个融合%s元素的游戏角色精灵图,正面站立姿势,轮廓清晰,适合作为2D游戏角色", tagStr) + case "background": + return fmt.Sprintf("一个%s风格的游戏场景背景,层次分明,包含前景、中景和远景", tagStr) + case "ui": + return fmt.Sprintf("一套%s风格的游戏UI元素,包括按钮、面板和图标", tagStr) + case "animation": + return fmt.Sprintf("一个%s风格的角色动画帧序列,动作流畅连贯", tagStr) + default: + return fmt.Sprintf("一个%s风格的游戏素材,高质量,适合2D游戏使用", tagStr) + } +} + +func buildStyle(tags []string) string { + parts := []string{"色彩鲜明", "细节丰富", "统一的光源方向"} + for _, t := range tags { + parts = append(parts, t+"风格") + } + return strings.Join(parts, ";") +} + +func buildTechNotes(assetType string) string { + switch assetType { + case "sprite": + return "输出格式: spritesheet;分辨率: 64x64 或 128x128;透明背景" + case "background": + return "输出格式: 独立PNG;分辨率: 1920x1080;层次分明的前中后景" + case "ui": + return "输出格式: 独立PNG素材;分辨率: 按元素适配;支持九宫格缩放" + case "animation": + return "输出格式: spritesheet或帧序列;建议4方向x4帧;透明背景" + default: + return "输出格式: PNG;分辨率: 标准2D游戏分辨率" + } +} diff --git a/backend/internal/service/prompt_agent_test.go b/backend/internal/service/prompt_agent_test.go new file mode 100644 index 0000000..f805286 --- /dev/null +++ b/backend/internal/service/prompt_agent_test.go @@ -0,0 +1,343 @@ +package service + +import ( + "context" + "net/http" + "net/http/httptest" + "strings" + "testing" +) + +func TestRunPromptAgent_Fallback(t *testing.T) { + // 无 API key 时走模板回退路径 + output, err := RunPromptAgent(context.Background(), PromptAgentInput{ + Tags: []string{"像素", "中世纪", "战士"}, + AssetType: "sprite", + UserNote: "持盾", + }) + if err != nil { + t.Fatalf("RunPromptAgent failed: %v", err) + } + + if output.Prompt == "" { + t.Fatal("expected non-empty prompt") + } + for _, section := range []string{"【主题】", "【风格】", "【技术】"} { + if !strings.Contains(output.Prompt, section) { + t.Errorf("output missing section %s: %s", section, output.Prompt) + } + } +} + +func TestRunPromptAgent_Sprite(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, "spritesheet") { + t.Errorf("sprite output should mention spritesheet: %s", output.Prompt) + } +} + +func TestRunPromptAgent_Background(t *testing.T) { + output, err := RunPromptAgent(context.Background(), PromptAgentInput{ + Tags: []string{"森林", "暗黑"}, + AssetType: "background", + }) + if err != nil { + t.Fatalf("RunPromptAgent failed: %v", err) + } + if !strings.Contains(output.Prompt, "1920x1080") { + t.Errorf("background output should mention 1920x1080: %s", output.Prompt) + } +} + +func TestRunPromptAgent_UI(t *testing.T) { + output, err := RunPromptAgent(context.Background(), PromptAgentInput{ + Tags: []string{"简约", "科幻"}, + AssetType: "ui", + }) + if err != nil { + t.Fatalf("RunPromptAgent failed: %v", err) + } + if !strings.Contains(output.Prompt, "九宫格") { + t.Errorf("ui output should mention 九宫格: %s", output.Prompt) + } +} + +func TestRunPromptAgent_Animation(t *testing.T) { + output, err := RunPromptAgent(context.Background(), PromptAgentInput{ + Tags: []string{"火焰", "魔法"}, + AssetType: "animation", + }) + if err != nil { + t.Fatalf("RunPromptAgent failed: %v", err) + } + if !strings.Contains(output.Prompt, "帧") { + t.Errorf("animation output should mention 帧: %s", output.Prompt) + } +} + +func TestRunPromptAgent_EmptyUserNote(t *testing.T) { + output, err := RunPromptAgent(context.Background(), PromptAgentInput{ + Tags: []string{"水彩"}, + AssetType: "sprite", + }) + if err != nil { + t.Fatalf("RunPromptAgent failed: %v", err) + } + if output.Prompt == "" { + t.Fatal("expected non-empty prompt") + } +} + +func TestRunPromptAgent_SingleTag(t *testing.T) { + output, err := RunPromptAgent(context.Background(), PromptAgentInput{ + Tags: []string{"赛博朋克"}, + AssetType: "background", + }) + if err != nil { + t.Fatalf("RunPromptAgent with single tag failed: %v", err) + } + if output.Prompt == "" { + t.Fatal("expected non-empty prompt") + } +} + +func TestRunPromptAgent_WithChatModel(t *testing.T) { + // 启动一个 mock OpenAI 兼容 API + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/chat/completions" { + http.NotFound(w, r) + return + } + w.Header().Set("Content-Type", "application/json") + w.Write([]byte(`{"choices":[{"message":{"role":"assistant","content":"【主题】一个像素战士\n【风格】像素风格\n【技术】spritesheet"}}]}`)) + })) + defer server.Close() + + // 保存原始配置,测试后恢复 + orig := llmCfg + defer func() { llmCfg = orig }() + llmCfg.BaseURL = server.URL + llmCfg.APIKey = "test-key" + llmCfg.Model = "test-model" + llmCfg.Temperature = 0.7 + llmCfg.MaxTokens = 512 + + 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("expected LLM response in output, got: %s", output.Prompt) + } + if output.RawText != output.Prompt { + t.Error("expected RawText == Prompt for non-streaming response") + } +} + +func TestChatCompletion_NonStream(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.Write([]byte(`{"choices":[{"message":{"content":"test response"}}]}`)) + })) + defer server.Close() + + orig := llmCfg + defer func() { llmCfg = orig }() + llmCfg.BaseURL = server.URL + llmCfg.APIKey = "key" + llmCfg.Model = "m" + llmCfg.Temperature = 0.5 + + result, err := chatCompletion(context.Background(), []chatMessage{ + {Role: "user", Content: "hello"}, + }) + if err != nil { + t.Fatalf("chatCompletion failed: %v", err) + } + if result != "test response" { + t.Errorf("expected 'test response', got %q", result) + } +} + +func TestChatCompletion_Stream(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + w.Write([]byte("data: {\"choices\":[{\"delta\":{\"content\":\"Hello\"}}]}\n\n")) + w.Write([]byte("data: {\"choices\":[{\"delta\":{\"content\":\" World\"}}]}\n\n")) + w.Write([]byte("data: [DONE]\n\n")) + })) + defer server.Close() + + orig := llmCfg + defer func() { llmCfg = orig }() + llmCfg.BaseURL = server.URL + llmCfg.APIKey = "key" + llmCfg.Model = "m" + llmCfg.Temperature = 0.5 + + result, err := chatCompletion(context.Background(), []chatMessage{ + {Role: "user", Content: "hi"}, + }) + if err != nil { + t.Fatalf("chatCompletion stream failed: %v", err) + } + if result != "Hello World" { + t.Errorf("expected 'Hello World', got %q", result) + } +} + +func TestCallLLMRefine_FallbackOnError(t *testing.T) { + orig := llmCfg + defer func() { llmCfg = orig }() + llmCfg.BaseURL = "http://invalid-url.invalid" + llmCfg.APIKey = "key" + llmCfg.Model = "m" + llmCfg.Temperature = 0.5 + + // API 调用失败时应回退到模板生成 + output, err := callLLMRefine(context.Background(), "用户选择标签: 测试\n素材类型: sprite") + if err != nil { + t.Fatalf("fallback should not error: %v", err) + } + if !strings.Contains(output.Prompt, "【主题】") { + t.Error("fallback output should contain 三段式 structure") + } +} + +// 单元测试 + +func TestBuildMetaPrompt(t *testing.T) { + in := PromptAgentInput{ + Tags: []string{"像素", "地牢"}, + AssetType: "sprite", + UserNote: "需要发光效果", + } + result := buildMetaPrompt(in) + + checks := []string{ + "2D 游戏素材提示词工程师", + "三段式结构", + "用户选择标签: 像素、地牢", + "素材类型: sprite", + "补充说明: 需要发光效果", + } + for _, c := range checks { + if !strings.Contains(result, c) { + t.Errorf("buildMetaPrompt missing %q", c) + } + } +} + +func TestBuildMetaPrompt_NoUserNote(t *testing.T) { + in := PromptAgentInput{ + Tags: []string{"像素"}, + AssetType: "sprite", + } + result := buildMetaPrompt(in) + if strings.Contains(result, "补充说明") { + t.Error("buildMetaPrompt should not contain 补充说明 when UserNote is empty") + } +} + +func TestParseTagsFromMeta(t *testing.T) { + meta := `用户选择标签: 像素、中世纪、战士 +素材类型: sprite` + tags, assetType := parseTagsFromMeta(meta) + if len(tags) != 3 { + t.Fatalf("expected 3 tags, got %d: %v", len(tags), tags) + } + if tags[0] != "像素" || tags[1] != "中世纪" || tags[2] != "战士" { + t.Errorf("unexpected tags: %v", tags) + } + if assetType != "sprite" { + t.Errorf("expected assetType=sprite, got %s", assetType) + } +} + +func TestParseTagsFromMeta_SingleTag(t *testing.T) { + meta := `用户选择标签: 赛博朋克 +素材类型: background` + 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") + 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") + if !strings.HasPrefix(prompt, "【主题】") { + t.Error("prompt should start with 【主题】") + } + themeIdx := strings.Index(prompt, "【主题】") + styleIdx := strings.Index(prompt, "【风格】") + techIdx := strings.Index(prompt, "【技术】") + if !(themeIdx < styleIdx && styleIdx < techIdx) { + t.Error("sections should be ordered: 主题 → 风格 → 技术") + } +} + +func TestBuildSubject(t *testing.T) { + tags := []string{"像素", "战士"} + tests := []struct { + assetType, want string + }{ + {"sprite", "精灵图"}, + {"background", "场景背景"}, + {"ui", "UI元素"}, + {"animation", "动画帧序列"}, + {"unknown", "游戏素材"}, + } + for _, tt := range tests { + result := buildSubject(tags, tt.assetType) + if !strings.Contains(result, tt.want) { + t.Errorf("buildSubject(%q) = %s, want containing %q", tt.assetType, result, tt.want) + } + } +} + +func TestBuildStyle(t *testing.T) { + result := buildStyle([]string{"像素", "暗黑"}) + if !strings.Contains(result, "色彩鲜明") { + t.Error("style should contain default descriptions") + } + for _, tag := range []string{"像素", "暗黑"} { + if !strings.Contains(result, tag+"风格") { + t.Errorf("style missing %q", tag+"风格") + } + } +} + +func TestBuildTechNotes(t *testing.T) { + tests := []struct { + assetType, want string + }{ + {"sprite", "spritesheet"}, + {"background", "1920x1080"}, + {"ui", "九宫格"}, + {"animation", "4方向x4帧"}, + {"unknown", "PNG"}, + } + for _, tt := range tests { + result := buildTechNotes(tt.assetType) + if !strings.Contains(result, tt.want) { + t.Errorf("buildTechNotes(%q) = %s, want containing %q", tt.assetType, result, tt.want) + } + } +}