feat(prompt-agent): 新增 PromptAgent Eino Chain 调用 LLM 优化提示词
- Eino Chain: formatMetaPrompt → llmRefine → 三段式输出 - callLLMRefine 调用 OpenAI 兼容 Chat Completions API - 支持 SSE 流式与非流式响应解析 - 无 API key 或调用失败时自动回退模板生成 - 新增完整的单元测试与集成测试(含 mock HTTP server)
This commit is contained in:
@@ -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游戏分辨率"
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user