aff48b6780
- buildMetaPrompt 按素材类型给出双模式指令(单素材/网格瓦片集),LLM 根据用户意图选择 - 新增 isSheetRequest 自动识别精灵表/瓦片集/tileset 关键字 - 模板回退链路全面支持 isSheet 双模式,默认单素材模式 - 所有素材类型统一纯白色背景(#FFFFFF),由后期 format 节点清洗去背 - 补充 sprite/background/ui/animation 四类素材的完整生成场景覆盖
385 lines
13 KiB
Go
Executable File
385 lines
13 KiB
Go
Executable File
package service
|
||
|
||
import (
|
||
"bufio"
|
||
"bytes"
|
||
"context"
|
||
"encoding/json"
|
||
"fmt"
|
||
"io"
|
||
"net/http"
|
||
"strings"
|
||
|
||
"gen2d/internal/config"
|
||
"gen2d/internal/logger"
|
||
|
||
"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")
|
||
sb.WriteString("4. 根据用户标签和描述,识别素材布局模式,在【技术】段明确标注格式:\n")
|
||
switch in.AssetType {
|
||
case "sprite":
|
||
sb.WriteString(" - 单个精灵(默认):独立PNG,纯白色背景(#FFFFFF),描述为单个角色立绘/道具图标\n")
|
||
sb.WriteString(" - 精灵表(spritesheet):角色/道具按行列等距网格排列,纯白色背景(#FFFFFF),帧间固定间距(2-4px),标注行列数,便于脚本一键拆分\n")
|
||
case "background":
|
||
sb.WriteString(" - 独立场景(默认):单张完整背景图,层次分明\n")
|
||
sb.WriteString(" - 场景瓦片集(tileset):地形/建筑元件按规则网格排列,纯白色背景(#FFFFFF),元件间固定间距,标注行列数与瓦片尺寸,确保无缝拼接\n")
|
||
case "ui":
|
||
sb.WriteString(" - 独立UI元素(默认):单个按钮/面板/图标,纯白色背景(#FFFFFF),独立PNG\n")
|
||
sb.WriteString(" - UI瓦片集(tileset):UI元件按规则网格排列,纯白色背景(#FFFFFF),元件间固定间距,标注行列数,支持九宫格缩放,便于脚本一键拆分\n")
|
||
case "animation":
|
||
sb.WriteString(" - 帧序列(默认):连续动画帧,独立帧文件或帧条带,纯白色背景(#FFFFFF)\n")
|
||
sb.WriteString(" - 动画精灵表(spritesheet):动画帧按行列等距网格排列,纯白色背景(#FFFFFF),帧间固定间距,标注行列数与方向数,便于脚本一键拆分\n")
|
||
}
|
||
sb.WriteString("\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) {
|
||
l := logger.FromCtx(ctx)
|
||
|
||
if llmCfg.APIKey == "" {
|
||
l.Warn("LLM API key not configured, using template fallback")
|
||
return fallbackRefine(metaPrompt), nil
|
||
}
|
||
|
||
messages := []chatMessage{
|
||
{Role: "system", Content: "你是一个专业的2D游戏素材提示词工程师,只输出优化后的中文提示词,不要任何解释。"},
|
||
{Role: "user", Content: metaPrompt},
|
||
}
|
||
|
||
l.Info("calling LLM API", "model", llmCfg.Model)
|
||
result, err := chatCompletion(ctx, messages)
|
||
if err != nil {
|
||
l.Error("LLM API call failed, using template fallback", "error", err)
|
||
return fallbackRefine(metaPrompt), nil
|
||
}
|
||
|
||
l.Info("LLM API succeeded", "response_length", len(result))
|
||
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)
|
||
|
||
l := logger.FromCtx(ctx)
|
||
|
||
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))
|
||
l.Error("LLM API error", "status", resp.StatusCode, "body", string(b))
|
||
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, userPrompt := parseTagsFromMeta(metaPrompt)
|
||
isSheet := isSheetRequest(tags, userPrompt)
|
||
prompt := generateStructuredPrompt(tags, assetType, isSheet)
|
||
return PromptAgentOutput{Prompt: prompt, RawText: prompt}
|
||
}
|
||
|
||
func parseTagsFromMeta(meta string) (tags []string, assetType string, userPrompt 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, "素材类型: "))
|
||
}
|
||
if strings.HasPrefix(line, "用户原始描述:") {
|
||
userPrompt = strings.TrimSpace(strings.TrimPrefix(line, "用户原始描述: "))
|
||
}
|
||
}
|
||
return
|
||
}
|
||
|
||
// isSheetRequest 从标签和提示词中识别是否为网格/瓦片集/精灵表模式。
|
||
func isSheetRequest(tags []string, prompt string) bool {
|
||
keywords := []string{"精灵表", "spritesheet", "瓦片集", "tileset", "tilemap"}
|
||
for _, t := range tags {
|
||
tLower := strings.ToLower(t)
|
||
for _, kw := range keywords {
|
||
if strings.Contains(tLower, kw) {
|
||
return true
|
||
}
|
||
}
|
||
}
|
||
promptLower := strings.ToLower(prompt)
|
||
for _, kw := range keywords {
|
||
if strings.Contains(promptLower, kw) {
|
||
return true
|
||
}
|
||
}
|
||
return false
|
||
}
|
||
|
||
func generateStructuredPrompt(tags []string, assetType string, isSheet bool) string {
|
||
var sb strings.Builder
|
||
sb.WriteString("【主题】")
|
||
sb.WriteString(buildSubject(tags, assetType, isSheet))
|
||
sb.WriteString("\n")
|
||
sb.WriteString("【风格】")
|
||
sb.WriteString(buildStyle(tags))
|
||
sb.WriteString("\n")
|
||
sb.WriteString("【技术】")
|
||
sb.WriteString(buildTechNotes(assetType, isSheet))
|
||
sb.WriteString("\n")
|
||
return sb.String()
|
||
}
|
||
|
||
func buildSubject(tags []string, assetType string, isSheet bool) string {
|
||
tagStr := strings.Join(tags, "、")
|
||
switch assetType {
|
||
case "sprite":
|
||
if isSheet {
|
||
return fmt.Sprintf("一个融合%s元素的游戏角色精灵表,角色按行列等距网格排列,轮廓清晰,适合作为2D游戏角色", tagStr)
|
||
}
|
||
return fmt.Sprintf("一个融合%s元素的游戏角色精灵图,正面站立姿势,轮廓清晰,适合作为2D游戏角色", tagStr)
|
||
case "background":
|
||
if isSheet {
|
||
return fmt.Sprintf("一套%s风格的游戏场景瓦片集,地形/建筑元件按规则网格排列,适合2D游戏地图拼接", tagStr)
|
||
}
|
||
return fmt.Sprintf("一个%s风格的游戏场景背景,层次分明,包含前景、中景和远景", tagStr)
|
||
case "ui":
|
||
if isSheet {
|
||
return fmt.Sprintf("一套%s风格的UI瓦片集,UI元件按规则网格排列,适合脚本一键拆分", tagStr)
|
||
}
|
||
return fmt.Sprintf("一套%s风格的游戏UI元素,包括按钮、面板和图标", tagStr)
|
||
case "animation":
|
||
if isSheet {
|
||
return fmt.Sprintf("一个%s风格的角色动画精灵表,动画帧按行列等距网格排列,动作流畅连贯", tagStr)
|
||
}
|
||
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, isSheet bool) string {
|
||
switch assetType {
|
||
case "sprite":
|
||
if isSheet {
|
||
return "输出格式: spritesheet;行列间隔完全相等、可脚本一键拆分对齐;纯白色背景(#FFFFFF,无渐变无噪点);帧间固定间距(2-4px);标注行列数"
|
||
}
|
||
return "输出格式: 独立PNG;纯白色背景(#FFFFFF);尺寸: 按角色比例适配"
|
||
case "background":
|
||
if isSheet {
|
||
return "输出格式: 场景瓦片集(tileset);规则网格排列;纯白色背景(#FFFFFF);瓦片间固定间距;标注行列数与瓦片尺寸;确保无缝拼接"
|
||
}
|
||
return "输出格式: 独立PNG;分辨率: 1920x1080;层次分明的前中后景"
|
||
case "ui":
|
||
if isSheet {
|
||
return "输出格式: UI瓦片集(tileset);规则网格排列;纯白色背景(#FFFFFF);元件间固定间距;标注行列数;支持九宫格缩放;可脚本一键拆分"
|
||
}
|
||
return "输出格式: 独立PNG素材;纯白色背景(#FFFFFF);分辨率: 按元素适配;支持九宫格缩放"
|
||
case "animation":
|
||
if isSheet {
|
||
return "输出格式: 动画精灵表(spritesheet);行列间隔相等;纯白色背景(#FFFFFF,无渐变无噪点);帧间固定间距;标注行列数与方向数;可脚本一键拆分"
|
||
}
|
||
return "输出格式: 帧序列或帧条带;独立帧文件;纯白色背景(#FFFFFF);建议4方向x4帧"
|
||
default:
|
||
return "输出格式: PNG;纯白色背景(#FFFFFF);分辨率: 标准2D游戏分辨率"
|
||
}
|
||
}
|