449 lines
12 KiB
Go
449 lines
12 KiB
Go
package main
|
||
|
||
import (
|
||
"context"
|
||
"encoding/json"
|
||
"flag"
|
||
"fmt"
|
||
"os"
|
||
"os/exec"
|
||
"path/filepath"
|
||
"strings"
|
||
"sync"
|
||
|
||
_ "github.com/gogf/gf/contrib/drivers/sqlite/v2"
|
||
|
||
"github.com/gogf/gf/v2/database/gdb"
|
||
"github.com/gogf/gf/v2/frame/g"
|
||
"github.com/gogf/gf/v2/os/glog"
|
||
|
||
"36wisdom/biz/consts"
|
||
)
|
||
|
||
// genasset:离线生成视觉素材(SVG→PNG 打包入客户端)与对比点评(回填数据库)。
|
||
// 用法:
|
||
//
|
||
// go run ./cmd/genasset --list # 打印生成清单
|
||
// go run ./cmd/genasset --only=comments # 只生成对比点评
|
||
// go run ./cmd/genasset --only=visuals --strategy=1 # 只生成第 1 计素材
|
||
// go run ./cmd/genasset --only=scripts # 生成剧本草稿(workspace/scripts/)
|
||
// go run ./cmd/genasset --import-scripts # 草稿导入回填 scene_node.script
|
||
// go run ./cmd/genasset --force # 覆盖已有产物
|
||
//
|
||
// 幂等:产物存在即跳过(--force 覆盖);点评按 feedback_pros 是否为空判定。
|
||
var (
|
||
ctx = context.Background()
|
||
genClient omlx
|
||
assetDir = "workspace/genasset" // SVG 源 + 转换缓存
|
||
staticDir = "ui-src/static/generated" // 打包入客户端的 PNG
|
||
)
|
||
|
||
type genFlags struct {
|
||
listOnly bool
|
||
only string // comments | visuals | scripts
|
||
strategy int // 0 = 全部计策
|
||
force bool
|
||
importScripts bool
|
||
}
|
||
|
||
type visualItem struct {
|
||
kind string // scene | char | prop | strategy
|
||
key string
|
||
id int64
|
||
name string
|
||
desc string
|
||
}
|
||
|
||
func main() {
|
||
fl := parseFlags()
|
||
genClient = omlx{
|
||
endpoint: g.Cfg().MustGet(ctx, "genasset.endpoint", "http://127.0.0.1:18080").String(),
|
||
apiKey: g.Cfg().MustGet(ctx, "genasset.api_key", "wenwu901").String(),
|
||
model: g.Cfg().MustGet(ctx, "genasset.model", "Qwen3.5-9B-MLX-4bit").String(),
|
||
timeout: g.Cfg().MustGet(ctx, "genasset.timeout", 600).Int(),
|
||
maxTokens: g.Cfg().MustGet(ctx, "genasset.max_tokens", 24576).Int(),
|
||
retries: g.Cfg().MustGet(ctx, "genasset.retries", 3).Int(),
|
||
}
|
||
if fl.listOnly {
|
||
printInventory(fl)
|
||
return
|
||
}
|
||
if fl.importScripts {
|
||
importScripts(fl)
|
||
return
|
||
}
|
||
if fl.only == "scripts" {
|
||
genScripts(fl)
|
||
return
|
||
}
|
||
if fl.only == "" || fl.only == "comments" {
|
||
genComments(fl)
|
||
}
|
||
if fl.only == "" || fl.only == "visuals" {
|
||
genVisuals(fl)
|
||
}
|
||
fmt.Println("[genasset] 完成")
|
||
}
|
||
|
||
func parseFlags() genFlags {
|
||
var fl genFlags
|
||
flag.BoolVar(&fl.listOnly, "list", false, "只打印生成清单")
|
||
flag.StringVar(&fl.only, "only", "", "comments | visuals,默认都生成")
|
||
flag.IntVar(&fl.strategy, "strategy", 0, "只处理指定计策(strategy_id),0=全部")
|
||
flag.BoolVar(&fl.force, "force", false, "覆盖已存在产物")
|
||
flag.BoolVar(&fl.importScripts, "import-scripts", false, "从草稿导入剧本到数据库(--strategy 限定计策)")
|
||
flag.Parse()
|
||
return fl
|
||
}
|
||
|
||
// ---------- 清单 ----------
|
||
|
||
func printInventory(fl genFlags) {
|
||
comments := pendingComments(fl)
|
||
fmt.Printf("待生成点评选项:%d(按节点分组 %d 个请求)\n", len(comments), distinctCount(comments))
|
||
fmt.Printf("待生成剧本节点:%d\n", len(pendingScriptNodes(fl)))
|
||
for _, it := range visualInventory(fl) {
|
||
fmt.Printf("%-9s %-16s %s\n", it.kind, it.key, it.name)
|
||
}
|
||
}
|
||
|
||
// ---------- 对比点评 ----------
|
||
|
||
func pendingComments(fl genFlags) []gdb.Record {
|
||
m := g.DB().Model(consts.TableNodeOption).Where("feedback_pros IS NULL OR feedback_pros = ''")
|
||
if fl.strategy > 0 {
|
||
_, nodeIDs := strategyScope(fl)
|
||
m = m.WhereIn("node_id", nodeIDs)
|
||
}
|
||
rows, err := m.Ctx(ctx).All()
|
||
if err != nil {
|
||
glog.Fatal(ctx, err)
|
||
}
|
||
return rows
|
||
}
|
||
|
||
func distinctCount(rows []gdb.Record) int {
|
||
seen := map[int64]bool{}
|
||
for _, r := range rows {
|
||
seen[r["node_id"].Int64()] = true
|
||
}
|
||
return len(seen)
|
||
}
|
||
|
||
// strategyScope 返回某计策涉及的 level_ids 与 node_ids(strategy<=0 时返回 nil)
|
||
func strategyScope(fl genFlags) ([]int64, []int64) {
|
||
if fl.strategy <= 0 {
|
||
return nil, nil
|
||
}
|
||
var levelIDs []int64
|
||
arr, err := g.DB().Model(consts.TableLevel).Where("strategy_id", fl.strategy).Array("id")
|
||
if err != nil {
|
||
glog.Fatal(ctx, err)
|
||
}
|
||
for _, v := range arr {
|
||
levelIDs = append(levelIDs, v.Int64())
|
||
}
|
||
var nodeIDs []int64
|
||
arr, err = g.DB().Model(consts.TableSceneNode).WhereIn("level_id", levelIDs).Array("id")
|
||
if err != nil {
|
||
glog.Fatal(ctx, err)
|
||
}
|
||
for _, v := range arr {
|
||
nodeIDs = append(nodeIDs, v.Int64())
|
||
}
|
||
return levelIDs, nodeIDs
|
||
}
|
||
|
||
func genComments(fl genFlags) {
|
||
rows := pendingComments(fl)
|
||
if len(rows) == 0 {
|
||
fmt.Println("[comments] 无待生成选项(全部已回填)")
|
||
return
|
||
}
|
||
byNode := map[int64][]gdb.Record{}
|
||
var nodeIDs []int64
|
||
for _, r := range rows {
|
||
nid := r["node_id"].Int64()
|
||
if _, ok := byNode[nid]; !ok {
|
||
nodeIDs = append(nodeIDs, nid)
|
||
}
|
||
byNode[nid] = append(byNode[nid], r)
|
||
}
|
||
nodeRows, err := g.DB().Model(consts.TableSceneNode).WhereIn("id", nodeIDs).Ctx(ctx).All()
|
||
if err != nil {
|
||
glog.Fatal(ctx, err)
|
||
}
|
||
nodes := map[int64]gdb.Record{}
|
||
for _, r := range nodeRows {
|
||
nodes[r["id"].Int64()] = r
|
||
}
|
||
sem := make(chan struct{}, 2) // 本地 9B 模型慢,并发 2 防排队打爆
|
||
var wg sync.WaitGroup
|
||
for _, nid := range nodeIDs {
|
||
wg.Add(1)
|
||
sem <- struct{}{}
|
||
go func(nid int64) {
|
||
defer wg.Done()
|
||
defer func() { <-sem }()
|
||
if err := genNodeComments(nid, nodes[nid], byNode[nid]); err != nil {
|
||
glog.Errorf(ctx, "[comments] node %d 失败: %v", nid, err)
|
||
}
|
||
}(nid)
|
||
}
|
||
wg.Wait()
|
||
}
|
||
|
||
func genNodeComments(nid int64, node gdb.Record, opts []gdb.Record) error {
|
||
var lines []string
|
||
for _, o := range opts {
|
||
lines = append(lines, fmt.Sprintf("%d. %s", o["id"].Int64(), o["text"].String()))
|
||
}
|
||
out, err := genClient.chat(ctx, commentSystem, fmt.Sprintf(commentUser, node["content"].String(), strings.Join(lines, "\n")))
|
||
if err != nil {
|
||
return err
|
||
}
|
||
var resp struct {
|
||
Comments []struct {
|
||
OptionID int64 `json:"option_id"`
|
||
Pros string `json:"pros"`
|
||
Cons string `json:"cons"`
|
||
} `json:"comments"`
|
||
}
|
||
if err := json.Unmarshal([]byte(out), &resp); err != nil {
|
||
return fmt.Errorf("点评 JSON 解析失败: %v(输出: %s)", err, truncate(out, 200))
|
||
}
|
||
if len(resp.Comments) == 0 {
|
||
return fmt.Errorf("点评空结果: %s", truncate(out, 200))
|
||
}
|
||
for _, c := range resp.Comments {
|
||
// 注意:Update 的变参是 dataAndWhere,直接 Update(ctx) 会把 ctx 当作 Data 覆盖掉,必须用 .Ctx(ctx)
|
||
_, err := g.DB().Model(consts.TableNodeOption).
|
||
Ctx(ctx).
|
||
Data(g.Map{"feedback_pros": c.Pros, "feedback_cons": c.Cons}).
|
||
Where("id", c.OptionID).Update()
|
||
if err != nil {
|
||
return err
|
||
}
|
||
}
|
||
fmt.Printf("[comments] ok node %d(%d 个选项)\n", nid, len(resp.Comments))
|
||
return nil
|
||
}
|
||
|
||
// ---------- 视觉素材 ----------
|
||
|
||
func visualInventory(fl genFlags) []visualItem {
|
||
var items []visualItem
|
||
items = append(items, elementItems("scene", sceneIDs(fl))...)
|
||
items = append(items, elementItems("char", charIDs(fl))...)
|
||
items = append(items, elementItems("prop", propIDs(fl))...)
|
||
// 计策卡图标
|
||
m := g.DB().Model(consts.TableStrategy)
|
||
if fl.strategy > 0 {
|
||
m = m.Where("id", fl.strategy)
|
||
}
|
||
rows, err := m.Ctx(ctx).All()
|
||
if err != nil {
|
||
glog.Fatal(ctx, err)
|
||
}
|
||
for _, r := range rows {
|
||
items = append(items, visualItem{
|
||
kind: "strategy", key: fmt.Sprintf("strategy_%d", r["id"].Int64()),
|
||
id: r["id"].Int64(), name: r["name"].String(), desc: r["meaning"].String(),
|
||
})
|
||
}
|
||
return items
|
||
}
|
||
|
||
func sceneIDs(fl genFlags) []int64 {
|
||
m := g.DB().Model(consts.TableLevel).Where("scene_id > 0").Fields("DISTINCT scene_id")
|
||
if fl.strategy > 0 {
|
||
m = m.Where("strategy_id", fl.strategy)
|
||
}
|
||
var ids []int64
|
||
arr, err := m.Array("scene_id")
|
||
if err != nil {
|
||
glog.Fatal(ctx, err)
|
||
}
|
||
for _, v := range arr {
|
||
ids = append(ids, v.Int64())
|
||
}
|
||
return ids
|
||
}
|
||
|
||
func charIDs(fl genFlags) []int64 {
|
||
m := g.DB().Model(consts.TableSceneNode).Where("character_id > 0").Fields("DISTINCT character_id")
|
||
if fl.strategy > 0 {
|
||
levelIDs, _ := strategyScope(fl)
|
||
m = m.WhereIn("level_id", levelIDs)
|
||
}
|
||
var ids []int64
|
||
arr, err := m.Array("character_id")
|
||
if err != nil {
|
||
glog.Fatal(ctx, err)
|
||
}
|
||
for _, v := range arr {
|
||
ids = append(ids, v.Int64())
|
||
}
|
||
return ids
|
||
}
|
||
|
||
func propIDs(fl genFlags) []int64 {
|
||
m := g.DB().Model(consts.TableNodeOption).Where("prop_id > 0").Fields("DISTINCT prop_id")
|
||
if fl.strategy > 0 {
|
||
_, nodeIDs := strategyScope(fl)
|
||
m = m.WhereIn("node_id", nodeIDs)
|
||
}
|
||
var ids []int64
|
||
arr, err := m.Array("prop_id")
|
||
if err != nil {
|
||
glog.Fatal(ctx, err)
|
||
}
|
||
for _, v := range arr {
|
||
ids = append(ids, v.Int64())
|
||
}
|
||
return ids
|
||
}
|
||
|
||
func elementItems(kind string, ids []int64) []visualItem {
|
||
if len(ids) == 0 {
|
||
return nil
|
||
}
|
||
rows, err := g.DB().Model(consts.TableElement).WhereIn("id", ids).Ctx(ctx).All()
|
||
if err != nil {
|
||
glog.Fatal(ctx, err)
|
||
}
|
||
var items []visualItem
|
||
for _, r := range rows {
|
||
items = append(items, visualItem{
|
||
kind: kind, key: fmt.Sprintf("%s_%d", kind, r["id"].Int64()),
|
||
id: r["id"].Int64(), name: r["name"].String(), desc: r["description"].String(),
|
||
})
|
||
}
|
||
return items
|
||
}
|
||
|
||
func genVisuals(fl genFlags) {
|
||
items := visualInventory(fl)
|
||
if len(items) == 0 {
|
||
fmt.Println("[visuals] 清单为空")
|
||
return
|
||
}
|
||
fmt.Printf("[visuals] 共 %d 个素材\n", len(items))
|
||
sem := make(chan struct{}, 2)
|
||
var wg sync.WaitGroup
|
||
for _, it := range items {
|
||
wg.Add(1)
|
||
sem <- struct{}{}
|
||
go func(it visualItem) {
|
||
defer wg.Done()
|
||
defer func() { <-sem }()
|
||
if err := genOneVisual(fl, it); err != nil {
|
||
glog.Errorf(ctx, "[visuals] %s %s 失败: %v", it.kind, it.key, err)
|
||
}
|
||
}(it)
|
||
}
|
||
wg.Wait()
|
||
// 统一转 PNG(每个 kind 一个目录)
|
||
for _, kind := range []string{"scene", "char", "prop", "strategy"} {
|
||
src := filepath.Join(assetDir, kind)
|
||
if _, err := os.Stat(src); err != nil {
|
||
continue
|
||
}
|
||
if err := svg2png(src, filepath.Join(staticDir, kind)); err != nil {
|
||
glog.Errorf(ctx, "[visuals] svg2png %s: %v", kind, err)
|
||
}
|
||
}
|
||
// 回填 image 路径(前端按 /static/... 直接引用);PNG 未生成成功的不回填(避免前端 404)
|
||
for _, it := range items {
|
||
if !fileExists(filepath.Join(staticDir, it.kind, it.key+".png")) {
|
||
glog.Errorf(ctx, "[visuals] 回填 %s 跳过:PNG 不存在", it.key)
|
||
continue
|
||
}
|
||
img := pngPath(it)
|
||
var err error
|
||
if it.kind == "strategy" {
|
||
_, err = g.DB().Model(consts.TableStrategy).Ctx(ctx).Data(g.Map{"icon": img}).Where("id", it.id).Update()
|
||
} else {
|
||
_, err = g.DB().Model(consts.TableElement).Ctx(ctx).Data(g.Map{"image": img}).Where("id", it.id).Update()
|
||
}
|
||
if err != nil {
|
||
glog.Errorf(ctx, "[visuals] 回填 %s 失败: %v", it.key, err)
|
||
}
|
||
}
|
||
}
|
||
|
||
func genOneVisual(fl genFlags, it visualItem) error {
|
||
svgPath := filepath.Join(assetDir, it.kind, it.key+".svg")
|
||
if !fl.force && fileExists(svgPath) {
|
||
return nil // 幂等:已有产物跳过
|
||
}
|
||
var user string
|
||
switch it.kind {
|
||
case "scene":
|
||
user = fmt.Sprintf("场景名称:%s\n场景说明:%s\n%s", it.name, it.desc, styleRule)
|
||
case "char":
|
||
user = fmt.Sprintf("角色名称:%s\n角色说明:%s\n%s", it.name, it.desc, styleRule)
|
||
case "prop":
|
||
user = fmt.Sprintf("道具名称:%s\n道具说明:%s\n%s", it.name, it.desc, styleRule)
|
||
case "strategy":
|
||
user = fmt.Sprintf("计策名称:%s\n计策释义:%s\n%s", it.name, it.desc, styleRule)
|
||
}
|
||
out, err := genClient.chat(ctx, "", user)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
svg, err := extractSVG(out)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
if err := os.MkdirAll(filepath.Dir(svgPath), 0o755); err != nil {
|
||
return err
|
||
}
|
||
if err := os.WriteFile(svgPath, []byte(svg), 0o644); err != nil {
|
||
return err
|
||
}
|
||
fmt.Printf("[visuals] ok %s/%s\n", it.kind, it.key)
|
||
return nil
|
||
}
|
||
|
||
func extractSVG(out string) (string, error) {
|
||
var box struct {
|
||
SVG string `json:"svg"`
|
||
}
|
||
if err := json.Unmarshal([]byte(out), &box); err != nil {
|
||
return "", fmt.Errorf("模型输出不是 JSON: %v(输出: %s)", err, truncate(out, 200))
|
||
}
|
||
if !strings.Contains(box.SVG, "<svg") {
|
||
return "", fmt.Errorf("输出缺少 svg 内容: %s", truncate(out, 200))
|
||
}
|
||
return box.SVG, nil
|
||
}
|
||
|
||
func pngPath(it visualItem) string {
|
||
return "/static/generated/" + it.kind + "/" + it.key + ".png"
|
||
}
|
||
|
||
// svg2png 调 ui-src/scripts/svg2png.mjs(sharp)批量转换
|
||
func svg2png(src, dst string) error {
|
||
cmd := exec.CommandContext(ctx, "node", "ui-src/scripts/svg2png.mjs", src, dst)
|
||
out, err := cmd.CombinedOutput()
|
||
if err != nil {
|
||
return fmt.Errorf("%s: %v", strings.TrimSpace(string(out)), err)
|
||
}
|
||
fmt.Print(string(out))
|
||
return nil
|
||
}
|
||
|
||
func fileExists(p string) bool {
|
||
_, err := os.Stat(p)
|
||
return err == nil
|
||
}
|
||
|
||
func truncate(s string, n int) string {
|
||
r := []rune(s)
|
||
if len(r) <= n {
|
||
return s
|
||
}
|
||
return string(r[:n]) + "…"
|
||
}
|