Files
36Wisdom/cmd/genasset/main.go
T

436 lines
12 KiB
Go
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 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 --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
strategy int // 0 = 全部计策
force 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.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.Parse()
return fl
}
// ---------- 清单 ----------
func printInventory(fl genFlags) {
comments := pendingComments(fl)
fmt.Printf("待生成点评选项:%d(按节点分组 %d 个请求)\n", len(comments), distinctCount(comments))
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_idsstrategy<=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.mjssharp)批量转换
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]) + "…"
}