Files
gen2d/backend/internal/service/prompt_agent.go
T
Gmarker689 ecb74a2aed feat: 异步生成管线 + 前端轮询 + 图片编辑 + JWT 中间件
后端:
- 异步生成: POST /api/v1/generate 立即返回 taskId,后台执行管线
- 任务轮询: GET /api/v1/tasks/:id + GET /api/v1/tasks/:id/assets
- 图片保存: 生成图片写入 ../generation/{projectId}/{taskId}/,静态服务
- 图片编辑: POST /api/v1/images/edit (multipart/form-data)
- JWT 中间件: mildware/auth.go 保护生成/编辑端点
- config.yml 清空敏感默认值,交由 .env 控制
- ImageGenConfig 新增 Quality 字段

前端:
- api/generate.ts: 对接真实 API (submitGenerate + poll getTask/getAssets)
- api/types.ts: 新增 GenerateResponse, AssetsResponse, Task 类型
- stores/generation.ts: 异步提交→轮询进度→获取素材→完成
- stores/task.ts: 默认分辨率 256→1024
- GenerateForm: 分辨率范围 1024-1536
- GeneratePage: 显示状态文本,完成后可查看结果/继续生成
- ResultPage: 从 store 读取,下载功能实现
2026-05-25 14:08:08 +08:00

314 lines
9.4 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 (
"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游戏分辨率"
}
}