#!/usr/bin/env bash
# gen2.sh — 即梦出图（可靠版）：提交 → 取 submit_id（stdout 优先，tasks.db 反查兜底）→ 下载 → 命名 → 记积分
#
# 为什么需要它：原 gen.sh 在**并发**下用 `list_task` 兜底取“最近一条任务”，会把别的 worker 的任务
# 当成自己的（实测出现多张完全相同的 PNG）。本脚本：
#   1) 提交用 `--poll=1`，CLI 会立刻打印含 submit_id 的 JSON（不再走长轮询的失败路径）；
#   2) 若 stdout 解析失败，则按**精确 Prompt 文本**在 ~/.dreamina_cli/tasks.db 的 aigc_task 表反查；
#   3) 下载轮询用 `query_result`（tasks.db 的 gen_status 只有被查询时才更新，不能只读库等）。
#
# 用法：
#   bash gen2.sh --name <语义名> --prompt-file <提示词文件> [--outdir <目录>]
#                [--ratio 16:9] [--model 5.0] [--res 2k] [--num 1]
set -uo pipefail
export PATH="$HOME/.local/bin:$PATH"

NAME=""; PROMPT_FILE=""; OUTDIR="."; RATIO="16:9"; MODEL="5.0"; RES="2k"; NUM=1
while [ $# -gt 0 ]; do
  case "$1" in
    --name)        NAME="$2"; shift 2;;
    --prompt-file) PROMPT_FILE="$2"; shift 2;;
    --outdir)      OUTDIR="$2"; shift 2;;
    --ratio)       RATIO="$2"; shift 2;;
    --model)       MODEL="$2"; shift 2;;
    --res)         RES="$2"; shift 2;;
    --num)         NUM="$2"; shift 2;;
    *) echo "未知参数: $1" >&2; exit 2;;
  esac
done
[ -n "$NAME" ] && [ -n "$PROMPT_FILE" ] || { echo "缺 --name / --prompt-file" >&2; exit 2; }
[ -f "$PROMPT_FILE" ] || { echo "找不到提示词文件: $PROMPT_FILE" >&2; exit 2; }
mkdir -p "$OUTDIR"
PROMPT="$(cat "$PROMPT_FILE")"

credit() { dreamina user_credit 2>/dev/null | python3 -c 'import sys,json,re
t=sys.stdin.read(); m=re.search(r"\{.*\}",t,re.S)
print(json.loads(m.group(0))["total_credit"] if m else "?")' 2>/dev/null || echo "?"; }

parse_sid() { python3 -c '
import sys, json, re
t = open(sys.argv[1], encoding="utf-8", errors="replace").read()
for m in re.finditer(r"\{.*?\}", t, re.S):
    try:
        d = json.loads(m.group(0))
    except Exception:
        continue
    if d.get("submit_id"):
        print(d["submit_id"]); break
' "$1" 2>/dev/null; }

db_sid() { # 按精确 Prompt 反查（唯一化，避免并发歧义）
python3 - "$1" "$2" <<'PY'
import sqlite3, json, sys, os
since, prompt = int(sys.argv[1]), sys.argv[2]
uri = "file:%s?mode=ro" % os.path.expanduser("~/.dreamina_cli/tasks.db")
db = sqlite3.connect(uri, uri=True)
for sid, req in db.execute("SELECT submit_id, request FROM aigc_task WHERE create_time >= ? ORDER BY create_time DESC", (since,)):
    try:
        body = json.loads(json.loads(req)["body"])
    except Exception:
        continue
    if body.get("Prompt") == prompt or body.get("prompt") == prompt:
        print(sid); break
PY
}

CRED_BEFORE=$(credit)
T0=$(( $(date +%s) - 3 ))
OUTLOG=/tmp/gen2_$$.out; ERRLOG=/tmp/gen2_$$.err

timeout 150 dreamina text2image --prompt="$PROMPT" --ratio="$RATIO" --resolution_type="$RES" \
        --model_version="$MODEL" --generate_num="$NUM" --poll=1 >"$OUTLOG" 2>"$ERRLOG"
CLI_RC=$?

SID=$(parse_sid "$OUTLOG")
SRC="stdout"
if [ -z "$SID" ]; then
  SID=$(db_sid "$T0" "$PROMPT"); SRC="tasks.db"
fi
if [ -z "$SID" ]; then
  echo "!! 未取得 submit_id（CLI rc=$CLI_RC）" >&2
  tail -5 "$ERRLOG" >&2; tail -5 "$OUTLOG" >&2
  exit 1
fi
echo "submit_id=$SID (via $SRC)"

TMPD=$(mktemp -d)
DL_OK=0
prev_size=0
for i in $(seq 1 60); do
  dreamina query_result --submit_id="$SID" --download_dir="$TMPD" >/dev/null 2>&1
  n=$(find "$TMPD" -type f \( -name '*.png' -o -name '*.jpg' -o -name '*.jpeg' \) | wc -l)
  if [ "$n" -gt 0 ]; then
    # 文件已出现，但要等它写完整：比对两次采样的大小，并用 PIL 校验可解码（防截断图）
    sz=$(du -sb "$TMPD" | cut -f1)
    if [ "$sz" = "$prev_size" ] && python3 - "$TMPD" <<'PY'
import sys, os
d = sys.argv[1]
for f in os.listdir(d):
    p = os.path.join(d, f)
    if not os.path.isfile(p):
        continue
    with open(p, "rb") as fh:
        head = fh.read(16)
        fh.seek(max(0, os.path.getsize(p) - 32))
        tail = fh.read()
    if head.startswith(b"\x89PNG"):
        if b"IEND" not in tail:
            sys.exit(1)
    elif head.startswith(b"\xff\xd8"):
        if not tail.rstrip(b"\x00").endswith(b"\xff\xd9"):
            sys.exit(1)
    else:
        sys.exit(1)
sys.exit(0)
PY
    then DL_OK=1; break; fi
    prev_size="$sz"
  fi
  st=$(python3 -c '
import sqlite3,sys,os
try:
    db=sqlite3.connect("file:%s?mode=ro" % os.path.expanduser("~/.dreamina_cli/tasks.db"),uri=True)
    r=db.execute("SELECT gen_status FROM aigc_task WHERE submit_id=?",(sys.argv[1],)).fetchone()
    print(r[0] if r else "?")
except Exception: print("?")
' "$SID" 2>/dev/null)
  case "$st" in failed|fail) echo "!! 任务失败 status=$st" >&2; exit 1;; esac
  sleep 5
done
[ "$DL_OK" = "1" ] || { echo "!! 超时未下载到图片 submit_id=$SID" >&2; exit 1; }

FILES=$(find "$TMPD" -type f \( -name '*.png' -o -name '*.jpg' -o -name '*.jpeg' \) | sort)
total=$(echo "$FILES" | wc -l); idx=0
for f in $FILES; do
  idx=$((idx+1))
  if [ "$total" -eq 1 ]; then dest="$OUTDIR/$NAME.png"; else dest="$OUTDIR/${NAME}_${idx}.png"; fi
  cp "$f" "$dest"; echo "OK $dest"
done
rm -rf "$TMPD"

CRED_AFTER=$(credit)
if [ "$CRED_BEFORE" != "?" ] && [ "$CRED_AFTER" != "?" ]; then
  echo "credit: $CRED_BEFORE -> $CRED_AFTER (扣 $((CRED_BEFORE-CRED_AFTER)))"
else
  echo "credit: $CRED_BEFORE -> $CRED_AFTER"
fi
echo "submit_id: $SID"
