#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""DSH 机视频生成流水线 Web 界面（端口 8898）—— 三步式。

步骤 1：上传参考图 → 点「👁️ 识图」→ 显示模型对图片的识别结果
步骤 2：根据识图结果填写描述 → 点「📝 生成提示词」→ 显示六段式 Ref2VA 提示词
步骤 3：确认提示词 → 点「🎬 生成视频」→ 提交 GPU 机 V18 工作流 → 视频回传下载中心

运行：vlm-env/bin/python web_app.py
"""
import json
import sys
import time
import uuid
import threading
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parent))
from v18_pipeline import run_pipeline, describe_image, build_prompt_from_vision  # noqa: E402

from flask import Flask, request, jsonify, render_template_string, send_from_directory

APP_DIR = Path(__file__).resolve().parent
UPLOAD_DIR = APP_DIR / "uploads"
UPLOAD_DIR.mkdir(exist_ok=True)

# 任务状态存储
TASKS = {}
TASKS_LOCK = threading.Lock()
# 步骤1/2 的异步任务
STEPS = {}
STEPS_LOCK = threading.Lock()

app = Flask(__name__)
app.config["MAX_CONTENT_LENGTH"] = 64 * 1024 * 1024  # 64MB

HTML = """<!DOCTYPE html>
<html lang="zh-CN">
<head>
<meta charset="utf-8">
<title>DSH 视频生成流水线（V18 Ref2VA · 三步式）</title>
<meta name="viewport" content="width=device-width, initial-scale=1">
<style>
:root { --bg:#0f1115; --card:#171b22; --border:#2a3040; --fg:#e8ecf3; --muted:#8b93a5; --accent:#4f8cff; --ok:#3ddc84; --err:#ff5d5d; }
* { box-sizing:border-box; margin:0; padding:0; }
body { background:var(--bg); color:var(--fg); font-family:-apple-system,"PingFang SC","Microsoft YaHei",sans-serif; padding:24px; }
.wrap { max-width:960px; margin:0 auto; }
h1 { font-size:22px; margin-bottom:4px; }
.sub { color:var(--muted); font-size:13px; margin-bottom:20px; }
.card { background:var(--card); border:1px solid var(--border); border-radius:12px; padding:20px; margin-bottom:16px; }
.step-head { display:flex; align-items:center; gap:10px; margin-bottom:12px; }
.step-num { width:28px; height:28px; border-radius:50%; background:var(--accent); color:#fff; display:flex; align-items:center; justify-content:center; font-weight:700; font-size:14px; flex-shrink:0; }
.step-title { font-size:16px; font-weight:600; }
label { display:block; font-size:13px; color:var(--muted); margin:14px 0 6px; }
input[type=file] { width:100%; padding:10px; background:#10141c; border:1px dashed var(--border); border-radius:8px; color:var(--fg); }
textarea { width:100%; min-height:110px; padding:10px; background:#10141c; border:1px solid var(--border); border-radius:8px; color:var(--fg); font-size:13px; resize:vertical; font-family:ui-monospace,Consolas,monospace; }
textarea.prompt { min-height:240px; }
.grid { display:grid; grid-template-columns:1fr 1fr; gap:12px; }
.grid3 { display:grid; grid-template-columns:1fr 1fr 1fr; gap:12px; }
input[type=number], select { width:100%; padding:8px; background:#10141c; border:1px solid var(--border); border-radius:8px; color:var(--fg); }
.btn { width:100%; padding:13px; border:none; border-radius:10px; color:#fff; font-size:15px; font-weight:600; cursor:pointer; margin-top:14px; }
.btn-vision { background:#6a5acd; }
.btn-prompt { background:#b8860b; }
.btn-video { background:var(--accent); }
.btn:disabled { opacity:.5; cursor:not-allowed; }
.hint { font-size:12px; color:var(--muted); margin-top:8px; line-height:1.6; }
/* 结果区：始终可见 */
.result-zone { margin-top:14px; border:1px solid var(--border); border-radius:10px; overflow:hidden; }
.result-zone .rz-head { display:flex; align-items:center; justify-content:space-between; padding:8px 12px; background:#1d2433; font-size:13px; font-weight:600; border-bottom:1px solid var(--border); }
.result-zone .rz-status { font-size:12px; font-weight:400; }
.rz-status.idle { color:var(--muted); }
.rz-status.running { color:#7fc4ff; }
.rz-status.done { color:var(--ok); }
.rz-status.error { color:var(--err); }
.rz-body { padding:12px; font-size:13px; line-height:1.7; white-space:pre-wrap; min-height:50px; }
.rz-body.idle { color:var(--muted); font-style:italic; }
.rz-body.running { color:#c8d2e6; }
.rz-body.done { color:var(--fg); }
.rz-body.error { color:var(--err); }
.rz-body a { color:var(--ok); }
.result video { width:100%; border-radius:8px; background:#000; margin-top:10px; }
.task-list { margin-top:8px; }
.task-item { padding:10px; border-bottom:1px solid var(--border); font-size:13px; display:flex; justify-content:space-between; align-items:center; }
.task-item a { color:var(--accent); text-decoration:none; }
.tag { font-size:11px; padding:2px 8px; border-radius:20px; }
.tag.wait { background:#3a3f4d; color:var(--muted); }
.tag.run { background:#1a3a5c; color:#7fc4ff; }
.tag.ok { background:#132a1e; color:var(--ok); }
.tag.err { background:#331517; color:var(--err); }
</style>
</head>
<body>
<div class="wrap">
<h1>🎬 DSH 视频生成流水线（三步式）</h1>
<div class="sub">① 识图 → ② 生成提示词 → ③ 生成视频。每步独立按钮，结果直接显示在各步骤下方区域。</div>

<!-- ============ 步骤 1：识图 ============ -->
<div class="card">
  <div class="step-head"><div class="step-num">1</div><div class="step-title">上传参考图 + 识图</div></div>
  <label>参考图片（第 1 张 = &lt;Picture 1&gt;，可多选，上传顺序即编号）</label>
  <input type="file" name="images" id="images" multiple accept="image/*">
  <button class="btn btn-vision" id="btnVision">👁️ 识图</button>
  <div class="result-zone">
    <div class="rz-head">👁️ 识图结果 <span class="rz-status idle" id="visionStatus">等待识图</span></div>
    <div class="rz-body idle" id="visionBody">点击上方「👁️ 识图」按钮，识别结果会显示在这里。</div>
  </div>
</div>

<!-- ============ 步骤 2：生成提示词 ============ -->
<div class="card">
  <div class="step-head"><div class="step-num">2</div><div class="step-title">填写描述 + 生成提示词</div></div>
  <label>描述（自己填写：画面内容、动作、运镜、台词等。识图结果仅供参考，不会自动填入）</label>
  <textarea name="requirement" id="requirement" placeholder="例：人物在古老的石桥上奔跑，镜头从背后跟随，背景是云雾缭绕的山峦和远处的古塔。人物穿着古装，手里拿着油纸伞。台词：「要迟到了！」"></textarea>
  <button class="btn btn-prompt" id="btnPrompt">📝 生成提示词</button>
  <div class="result-zone">
    <div class="rz-head">📝 提示词结果 <span class="rz-status idle" id="promptStatus">等待生成</span></div>
    <div class="rz-body idle" id="promptBody">点击上方「📝 生成提示词」按钮，六段式提示词会显示在这里。</div>
  </div>
  <label>生成的提示词（可手动修改后再生成视频）</label>
  <textarea class="prompt" id="promptOut" placeholder="提示词生成后显示在这里，可编辑"></textarea>
</div>

<!-- ============ 步骤 3：生成视频 ============ -->
<div class="card">
  <div class="step-head"><div class="step-num">3</div><div class="step-title">生成视频</div></div>
  <div class="grid">
    <div>
      <label>时长（秒）</label>
      <input type="number" name="duration" id="duration" value="10" min="4" max="15">
    </div>
    <div>
      <label>分辨率</label>
      <select name="resolution" id="resolution">
        <option value="736x1280" selected>736×1280（竖屏 9:16）</option>
        <option value="832x1216">832×1216</option>
        <option value="1216x832">1216×832（横屏 16:9）</option>
        <option value="1280x736">1280×736</option>
        <option value="864x640">864×640</option>
      </select>
    </div>
  </div>
  <div class="grid3">
    <div>
      <label>步数</label>
      <input type="number" name="steps" id="steps" value="8" min="4" max="30">
    </div>
    <div>
      <label>Seed（0=随机）</label>
      <input type="number" name="seed" id="seed" value="666666" min="0">
    </div>
    <div>
      <label>帧率</label>
      <input type="number" name="fps" id="fps" value="24" min="8" max="60">
    </div>
  </div>
  <button class="btn btn-video" id="btnVideo">🎬 生成视频</button>
  <div class="result-zone">
    <div class="rz-head">🎬 视频结果 <span class="rz-status idle" id="videoStatus">等待生成</span></div>
    <div class="rz-body idle" id="videoBody">点击上方「🎬 生成视频」按钮，提交 GPU 机 V18 工作流，完成后视频显示在这里。</div>
  </div>
  <div class="hint">提示：GPU 机队列忙时需等待；单次生成约 2–15 分钟（视排队）。若未生成提示词，将自动走完整流程。</div>
</div>

<!-- ============ 任务列表 ============ -->
<div class="card">
  <h2 style="font-size:16px;margin-bottom:10px;">📋 最近任务</h2>
  <div class="task-list" id="tasklist">加载中...</div>
</div>
</div>

<script>
const $ = id => document.getElementById(id);
const fmt = t => { const d = new Date(t); return d.toLocaleString('zh-CN', {hour12:false}); };
function escapeHtml(s) { return (s||'').replace(/&/g,'&amp;').replace(/</g,'&lt;').replace(/>/g,'&gt;').replace(/"/g,'&quot;'); }

// 结果区状态更新：head 状态标签 + body 内容
function setZone(statusId, bodyId, state, html) {
  const st = $(statusId); st.className = 'rz-status ' + state;
  const bd = $(bodyId); bd.className = 'rz-body ' + state; bd.innerHTML = html;
}
const STATE_LABEL = { idle:'等待', running:'处理中...', done:'完成', error:'失败' };

let currentVision = '';
function getFormData() {
  const fd = new FormData();
  const files = $('images').files;
  for (const f of files) fd.append('images', f);
  fd.append('requirement', $('requirement').value);
  fd.append('duration', $('duration').value);
  fd.append('resolution', $('resolution').value);
  fd.append('steps', $('steps').value);
  fd.append('seed', $('seed').value);
  fd.append('fps', $('fps').value);
  return fd;
}

// ---------- 步骤 1：识图 ----------
$('btnVision').addEventListener('click', async () => {
  const btn = $('btnVision'); btn.disabled = true; btn.textContent = '⏳ 识图中...';
  setZone('visionStatus','visionBody','running','正在识别图片内容（约 1–2 分钟）...');
  const fd = getFormData();
  try {
    const r = await fetch('/api/vision', { method:'POST', body: fd });
    const d = await r.json();
    if (!r.ok) throw new Error(d.error || '识图失败');
    setZone('visionStatus','visionBody','running','✅ 任务已提交，识别中（约 1–2 分钟）...');
    pollStep('vision', d.job_id, btn, '👁️ 识图');
  } catch(err) {
    setZone('visionStatus','visionBody','error','❌ 识图发起失败: ' + err.message);
    btn.disabled = false; btn.textContent = '👁️ 识图';
  }
});

// ---------- 步骤 2：生成提示词 ----------
$('btnPrompt').addEventListener('click', async () => {
  const btn = $('btnPrompt'); btn.disabled = true; btn.textContent = '⏳ 生成中...';
  setZone('promptStatus','promptBody','running','正在生成六段式提示词（约 1–3 分钟）...');
  const fd = getFormData();
  if (currentVision) fd.append('vision', currentVision);
  try {
    const r = await fetch('/api/build-prompt', { method:'POST', body: fd });
    const d = await r.json();
    if (!r.ok) throw new Error(d.error || '生成失败');
    setZone('promptStatus','promptBody','running','✅ 任务已提交，生成中（约 1–3 分钟）...');
    pollStep('prompt', d.job_id, btn, '📝 生成提示词');
  } catch(err) {
    setZone('promptStatus','promptBody','error','❌ 发起失败: ' + err.message);
    btn.disabled = false; btn.textContent = '📝 生成提示词';
  }
});

// ---------- 步骤 3：生成视频 ----------
$('btnVideo').addEventListener('click', async () => {
  const btn = $('btnVideo'); btn.disabled = true; btn.textContent = '⏳ 提交中...';
  setZone('videoStatus','videoBody','running','提交 GPU 机 V18 工作流...');
  const fd = getFormData();
  fd.append('prompt_override', $('promptOut').value);
  if (currentVision) fd.append('vision_override', currentVision);
  try {
    const r = await fetch('/api/submit', { method:'POST', body: fd });
    const d = await r.json();
    if (!r.ok) throw new Error(d.error || '提交失败');
    setZone('videoStatus','videoBody','running','✅ 已提交 GPU 队列，生成中（约 2–15 分钟，视排队）...');
    pollVideo(d.task_id);
  } catch(err) {
    setZone('videoStatus','videoBody','error','❌ 提交失败: ' + err.message);
    btn.disabled = false; btn.textContent = '🎬 生成视频';
  }
});

// ---------- 轮询：步骤1/2 ----------
let stepTimers = {};
function pollStep(kind, jobId, btn, btnLabel) {
  if (stepTimers[kind]) clearInterval(stepTimers[kind]);
  stepTimers[kind] = setInterval(async () => {
    try {
      const r = await fetch('/api/steps/' + jobId);
      const j = await r.json();
      if (j.status === 'running') {
        // 保持 running 状态
      } else {
        clearInterval(stepTimers[kind]); stepTimers[kind] = null;
        btn.disabled = false; btn.textContent = btnLabel;
        if (j.status === 'done') {
          if (kind === 'vision') {
            currentVision = j.result || '';
            setZone('visionStatus','visionBody','done', escapeHtml(currentVision));
          } else {
            $('promptOut').value = j.result || '';
            setZone('promptStatus','promptBody','done', '六段式提示词已生成，见下方可编辑框。');
          }
        } else {
          const zid = kind === 'vision' ? 'vision' : 'prompt';
          setZone(zid+'Status', zid+'Body', 'error', '❌ ' + (j.error || '处理失败'));
        }
      }
    } catch(e) { /* ignore */ }
  }, 3000);
}

// ---------- 轮询：步骤3（视频任务） ----------
let videoTimer = null;
function pollVideo(taskId) {
  if (videoTimer) clearInterval(videoTimer);
  videoTimer = setInterval(async () => {
    try {
      const r = await fetch('/api/task/' + taskId);
      const t = await r.json();
      if (t.status === 'running' || t.status === 'queued') {
        setZone('videoStatus','videoBody','running', t.log || '生成中...');
      } else {
        clearInterval(videoTimer); videoTimer = null;
        $('btnVideo').disabled = false; $('btnVideo').textContent = '🎬 生成视频';
        if (t.status === 'done') {
          setZone('videoStatus','videoBody','done',
            '<a href="' + t.video_url + '" target="_blank">📥 打开视频（下载中心）</a>' +
            '<div class="result"><video controls src="' + t.video_url + '"></video></div>');
          refreshTasks();
        } else {
          setZone('videoStatus','videoBody','error','❌ 失败: ' + (t.error || '未知错误'));
        }
      }
    } catch(e) { /* ignore */ }
  }, 3000);
}

// ---------- 任务列表 ----------
async function refreshTasks() {
  try {
    const r = await fetch('/api/tasks');
    const tasks = await r.json();
    const el = $('tasklist');
    if (!tasks.length) { el.innerHTML = '<div class="hint">暂无任务</div>'; return; }
    el.innerHTML = tasks.map(t => {
      const tag = t.status === 'done' ? '<span class="tag ok">完成</span>'
        : t.status === 'error' ? '<span class="tag err">失败</span>'
        : t.status === 'running' ? '<span class="tag run">运行中</span>'
        : '<span class="tag wait">排队</span>';
      const link = t.video_url ? ` <a href="${t.video_url}" target="_blank">视频</a>` : '';
      return `<div class="task-item"><span>${fmt(t.created)} · ${(t.requirement||'').slice(0,30)}</span><span>${tag}${link}</span></div>`;
    }).join('');
  } catch(e) { $('tasklist').innerHTML = '<div class="hint">加载失败</div>'; }
}
refreshTasks();
setInterval(refreshTasks, 10000);
</script>
</body>
</html>
"""


def get_resolution(res_str: str):
    w, h = res_str.lower().split("x")
    return int(w), int(h)


def save_uploads(files, subdir: str):
    """保存上传文件，返回本地路径列表。"""
    job_dir = UPLOAD_DIR / subdir
    job_dir.mkdir(exist_ok=True)
    saved = []
    for i, f in enumerate(files):
        ext = Path(f.filename).suffix or ".jpg"
        p = job_dir / f"pic_{i+1}{ext}"
        f.save(p)
        saved.append(str(p))
    return saved


@app.route("/")
def index():
    return render_template_string(HTML)


# ---------------- 步骤 1：识图 ----------------
@app.route("/api/vision", methods=["POST"])
def vision():
    files = request.files.getlist("images")
    if not files:
        return jsonify({"error": "请先选择参考图片"}), 400
    job_id = uuid.uuid4().hex[:10]
    saved = save_uploads(files, f"vision_{job_id}")
    with STEPS_LOCK:
        STEPS[job_id] = {
            "job_id": job_id, "kind": "vision", "images": saved,
            "status": "running", "log": "正在识图...", "result": None,
            "created": time.time() * 1000, "error": None,
        }
    threading.Thread(target=vision_worker, args=(job_id,), daemon=True).start()
    return jsonify({"job_id": job_id})


def vision_worker(job_id: str):
    def set_job(**kw):
        with STEPS_LOCK:
            if job_id in STEPS:
                STEPS[job_id].update(kw)
    try:
        j = STEPS[job_id]
        set_job(log="正在识图（本地 Qwen2.5-VL-3B，约 1–2 分钟）...")
        vision = describe_image(j["images"][0])
        set_job(status="done", log="✅ 识图完成", result=vision)
    except Exception as e:
        import traceback
        traceback.print_exc()
        set_job(status="error", error=str(e), log=f"❌ {e}")


# ---------------- 步骤 2：生成提示词 ----------------
@app.route("/api/build-prompt", methods=["POST"])
def build_prompt():
    files = request.files.getlist("images")
    requirement = (request.form.get("requirement") or "").strip()
    if not files and not requirement:
        return jsonify({"error": "请选择参考图片或填写描述"}), 400
    job_id = uuid.uuid4().hex[:10]
    saved = save_uploads(files, f"prompt_{job_id}")
    vision = (request.form.get("vision") or "").strip()
    with STEPS_LOCK:
        STEPS[job_id] = {
            "job_id": job_id, "kind": "prompt", "images": saved,
            "requirement": requirement, "vision": vision,
            "status": "running", "log": "正在生成提示词...", "result": None,
            "created": time.time() * 1000, "error": None,
        }
    threading.Thread(target=prompt_worker, args=(job_id,), daemon=True).start()
    return jsonify({"job_id": job_id})


def prompt_worker(job_id: str):
    def set_job(**kw):
        with STEPS_LOCK:
            if job_id in STEPS:
                STEPS[job_id].update(kw)
    try:
        j = STEPS[job_id]
        set_job(log="正在生成六段式提示词（约 1–3 分钟）...")
        if j["images"]:
            prompt = build_prompt_from_vision(j["images"][0], j["vision"], j["requirement"], 10)
        else:
            # 无图：直接让模型按要求生成提示词
            prompt = build_prompt_from_vision(None, j["vision"], j["requirement"], 10)
        set_job(status="done", log="✅ 提示词生成完成", result=prompt)
    except Exception as e:
        import traceback
        traceback.print_exc()
        set_job(status="error", error=str(e), log=f"❌ {e}")


# ---------------- 步骤 3：生成视频 ----------------
@app.route("/api/submit", methods=["POST"])
def submit():
    files = request.files.getlist("images")
    prompt_override = (request.form.get("prompt_override") or "").strip()
    vision_override = (request.form.get("vision_override") or "").strip()
    requirement = (request.form.get("requirement") or "").strip()
    if not files:
        return jsonify({"error": "请选择参考图片"}), 400
    if not prompt_override and not requirement:
        return jsonify({"error": "请先在第 2 步生成提示词，或填写描述"}), 400

    task_id = uuid.uuid4().hex[:10]
    saved = save_uploads(files, task_id)

    duration = int(request.form.get("duration", 10))
    w, h = get_resolution(request.form.get("resolution", "736x1280"))
    steps = int(request.form.get("steps", 8))
    seed = int(request.form.get("seed", 666666))
    fps = float(request.form.get("fps", 24))
    if seed == 0:
        import random
        seed = random.randint(1, 2**31)

    with TASKS_LOCK:
        TASKS[task_id] = {
            "task_id": task_id,
            "images": saved,
            "requirement": requirement,
            "params": {"width": w, "height": h, "duration": duration, "steps": steps, "seed": seed, "fps": fps},
            "prompt_override": prompt_override,
            "vision_override": vision_override,
            "status": "queued",
            "created": time.time() * 1000,
            "log": "排队中...",
            "prompt": None,
            "vision_description": None,
            "video_url": None,
            "error": None,
        }

    threading.Thread(target=worker, args=(task_id,), daemon=True).start()
    return jsonify({"task_id": task_id})


def worker(task_id: str):
    def set_task(**kw):
        with TASKS_LOCK:
            if task_id in TASKS:
                TASKS[task_id].update(kw)

    try:
        t = TASKS[task_id]
        if t["prompt_override"]:
            set_task(status="running", log="使用你的提示词，提交 GPU 机（约 2–15 分钟）...")
        else:
            set_task(status="running", log="正在解析图片（识图阶段，约 1–2 分钟）...")
        result = run_pipeline(
            t["images"][0],
            t["requirement"],
            width=t["params"]["width"],
            height=t["params"]["height"],
            duration=t["params"]["duration"],
            seed=t["params"]["seed"],
            steps=t["params"]["steps"],
            fps=t["params"]["fps"],
            prompt_override=t["prompt_override"] or None,
            vision_override=t["vision_override"] or None,
            on_vision=lambda v: set_task(
                vision_description=v,
                log="✅ 识图完成，正在生成提示词并提交 GPU（约 5–15 分钟）..."),
        )
        set_task(status="done", log="✅ 完成", prompt=result["prompt"],
                 vision_description=result.get("vision_description", ""),
                 video_url=result["video_url"])
    except Exception as e:
        import traceback
        traceback.print_exc()
        set_task(status="error", error=str(e), log=f"❌ {e}")


@app.route("/api/task/<task_id>")
def task_status(task_id):
    with TASKS_LOCK:
        t = TASKS.get(task_id)
        if not t:
            return jsonify({"error": "not found"}), 404
        return jsonify(t)


@app.route("/api/tasks")
def task_list():
    with TASKS_LOCK:
        items = sorted(TASKS.values(), key=lambda x: x["created"], reverse=True)[:20]
        return jsonify(items)


@app.route("/api/steps/<job_id>")
def step_status(job_id):
    with STEPS_LOCK:
        j = STEPS.get(job_id)
        if not j:
            return jsonify({"error": "not found"}), 404
        return jsonify(j)


@app.route("/uploads/<path:filename>")
def uploaded_file(filename):
    return send_from_directory(str(UPLOAD_DIR), filename)


if __name__ == "__main__":
    print("DSH 视频流水线 Web（三步式）: http://192.168.31.76:8898")
    app.run(host="0.0.0.0", port=8898, debug=False, threaded=True)
