#!/usr/bin/env python3
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: © 2025 Tenstorrent AI ULC
"""
tt-ctl — lightweight CLI for tt-local-generator.

Works without the GUI running.  Reads and writes the same queue.json /
history.json files that the GUI persists on every mutation.

Usage:
    tt-ctl status                   Server health, queue depth, history summary
    tt-ctl servers                  Live status of every managed service
    tt-ctl start <service>          Start a service (wan2.2 | mochi | flux | animate |
                                    skyreels | artgen-qwen3-8b | prompt-server | all)
    tt-ctl start --single-chip      Start artgen + prompt-server only
                                    (single Blackhole card or CPU-only)
    tt-ctl stop  <service>          Stop a service (same keys as start)
    tt-ctl restart <service>        Stop then re-start a service
    tt-ctl queue                    List pending queue items
    tt-ctl queue add "prompt"       Append a prompt to the persistent queue
    tt-ctl queue clear              Empty the queue (does not cancel in-flight jobs)
    tt-ctl history [N]              Show last N completed generations (default 10)
    tt-ctl recover                  List server jobs not in local history
    tt-ctl run "prompt"             Run one generation right now (blocking)
                                    Flags: --steps N  --model wan2|mochi|flux
                                           --server URL  --neg "negative prompt"
    tt-ctl server                   Show server model / health detail
    tt-ctl logs [service]           Tail the log for a managed service
                                    Default: the running inference server
                                    Services: wan2.2 | mochi | flux | animate |
                                              skyreels | prompt-server | docker
                                    Flags: --lines N  --follow / -f
    tt-ctl serve-inventory          Serve the local video library over HTTP
                                    so a remote GUI (--server http://this:8000)
                                    can browse and download it.
                                    Flags: --port 8002  --host 0.0.0.0
    tt-ctl artgen                   Generate generative art via LLM — SVG landscapes,
                                    city skylines, constellations, geometric patterns,
                                    ANSI art, color palettes, verse, circuit diagrams.
                                    Flags: --type TYPE  --output PATH  --model ID
                                           --simulate  --freeform "describe anything"
                                    Types: landscape skyline constellation geometric
                                           ansi palette verse circuit freeform
    tt-ctl plugin list              List all discovered plugins (name, tab, output type)
    tt-ctl mcp-config               Emit Claude Code MCP config JSON for tt-local-gen
                                    (pipe to ~/.claude/mcp.json to register with Claude)

All commands accept --server URL to override the default http://localhost:8000.
"""

import argparse
import glob as _glob
import json
import os
import subprocess
import sys
import textwrap
import threading
from datetime import datetime
from pathlib import Path

# ── Make sure app/ is importable ─────────────────────────────────────────────
_HERE = Path(__file__).resolve().parent
sys.path.insert(0, str(_HERE / "app"))

try:
    _VERSION = (_HERE / "VERSION").read_text().strip()
except FileNotFoundError:
    _VERSION = "unknown"

from api_client import APIClient          # noqa: E402
from history_store import HistoryStore    # noqa: E402
from worker import (                                                    # noqa: E402
    GenerationWorker,
    ImageGenerationWorker,
    AnimateGenerationWorker,
)
import generate_prompt as gp                                            # noqa: E402
import server_manager as sm              # noqa: E402
from artgen.cli import cmd_artgen, _build_artgen_parser                 # noqa: E402


# ── Plugin management commands ────────────────────────────────────────────────

def cmd_plugin(args):
    """Dispatch plugin sub-commands."""
    if getattr(args, "plugin_cmd", None) == "list":
        _cmd_plugin_list(args)
    else:
        print("Usage: tt-ctl plugin list")


def _cmd_plugin_list(_args):
    """List all loaded plugins with name, media type, tab, hardware, and tools."""
    import plugin_loader
    plugin_loader.load_plugins()
    plugins = plugin_loader.all_plugins()
    if not plugins:
        print("No plugins loaded.")
        return
    print(f"{'NAME':<20} {'MEDIA':<10} {'TAB':<16} {'HARDWARE':<12} TOOLS")
    print("-" * 72)
    for p in plugins:
        xttlg = p.manifest.get("x-ttlg", {})
        hw = xttlg.get("hardware") or "—"
        media = xttlg.get("media_type", "—")
        tab = xttlg.get("tab", "—")
        tool_names = ", ".join(t["name"] for t in p.tools)
        print(f"{p.name:<20} {media:<10} {tab:<16} {hw:<12} {tool_names}")


def cmd_workflow(args):
    """Run or list generative workflow specs."""
    if not hasattr(args, "workflow_cmd") or args.workflow_cmd is None:
        print("Usage: tt-ctl workflow run <spec.json> [--dry-run]")
        print("       tt-ctl workflow list")
        return

    if args.workflow_cmd == "list":
        examples_dir = _HERE / "docs" / "examples" / "workflows"
        if not examples_dir.exists():
            print("No workflow examples found.")
            return
        print("Available workflows:")
        for f in sorted(examples_dir.glob("*.json")):
            import json as _json
            try:
                meta = _json.loads(f.read_text())
                desc = meta.get("_description", "")[:80]
            except Exception:
                desc = ""
            print(f"  {f.name}")
            if desc:
                print(f"    {desc}")
        return

    if args.workflow_cmd == "run":
        spec_path = args.spec
        dry_run = getattr(args, "dry_run", False)
        runner = _HERE / "bin" / "run_workflow.sh"
        if not runner.exists():
            print(f"ERROR: {runner} not found", file=sys.stderr)
            sys.exit(1)
        cmd = ["bash", str(runner), spec_path]
        if dry_run:
            cmd.append("--dry-run")
        import subprocess as _sp
        result = _sp.run(cmd)
        sys.exit(result.returncode)


def cmd_mcp_config(_args):
    """Emit Claude Code MCP config JSON for tt-local-gen."""
    import json as _json
    port = int(os.environ.get("TTLG_MCP_PORT", "8120"))
    config = {"tt-local-gen": {"url": f"http://localhost:{port}/mcp"}}
    print(_json.dumps(config, indent=2))
    # Guidance to stderr so stdout stays machine-readable.
    # WARNING: appending with >> will corrupt an existing JSON file.
    # Merge the key manually, or use this for a fresh config:
    mcp_url = f"http://localhost:{port}/mcp"
    print("\n# To add to Claude Code (merge into existing ~/.claude/mcp.json):", file=sys.stderr)
    print("# WARNING: do not use >> — that corrupts existing JSON. Merge instead:", file=sys.stderr)
    print("#   python3 -c '", file=sys.stderr)
    print("#     import json, pathlib", file=sys.stderr)
    print("#     p = pathlib.Path.home() / \".claude/mcp.json\"", file=sys.stderr)
    print("#     cfg = json.loads(p.read_text()) if p.exists() else {}", file=sys.stderr)
    print(f'#     cfg["tt-local-gen"] = {{"url": "{mcp_url}"}}', file=sys.stderr)
    print("#     p.write_text(json.dumps(cfg, indent=2))", file=sys.stderr)
    print("#   '", file=sys.stderr)


# ── Colour helpers (plain text when stdout is not a tty) ──────────────────────

def _tty():
    return sys.stdout.isatty()

def _c(code, text):
    return f"\033[{code}m{text}\033[0m" if _tty() else text

def green(t):  return _c("32", t)
def yellow(t): return _c("33", t)
def red(t):    return _c("31", t)
def teal(t):   return _c("36", t)
def bold(t):   return _c("1",  t)
def dim(t):    return _c("2",  t)


# ── Helpers ───────────────────────────────────────────────────────────────────

def _make_client(args) -> APIClient:
    url = getattr(args, "server", "http://localhost:8000")
    return APIClient(base_url=url)


def _wrap(text: str, width: int = 80, indent: str = "  ") -> str:
    return textwrap.fill(text, width=width, initial_indent=indent,
                         subsequent_indent=indent)


def _age(iso: str) -> str:
    """Human-readable age from an ISO 8601 timestamp."""
    try:
        dt = datetime.fromisoformat(iso)
        delta = datetime.now() - dt
        s = int(delta.total_seconds())
        if s < 60:
            return f"{s}s ago"
        if s < 3600:
            return f"{s//60}m ago"
        if s < 86400:
            return f"{s//3600}h ago"
        return f"{s//86400}d ago"
    except Exception:
        return ""


# ── Worker selection ──────────────────────────────────────────────────────────

# Maps model_source values stored in queue items to the model identifier string
# forwarded to the inference server.
_MODEL_FOR_SRC: dict[str, str] = {
    "video":    "wan2.2-t2v",
    "mochi":    "mochi-1-preview",
    "skyreels": "skyreels-v2-df",
    "image":    "flux.1-dev",
    "animate":  "wan2.2-animate-14b",
}


def _make_worker_for_item(client: "APIClient", store: "HistoryStore", item: dict):
    """
    Instantiate the correct generation worker for a queue item dict.

    Selects GenerationWorker, ImageGenerationWorker, or AnimateGenerationWorker
    based on item["model_source"].  Unknown model_source values cause cmd_queue_run
    to warn and skip the item (leaving it in the queue).
    """
    src    = item.get("model_source", "video")
    model  = _MODEL_FOR_SRC.get(src, "wan2.2-t2v")
    prompt = item.get("prompt", "")
    neg    = item.get("negative_prompt", "")
    steps  = item.get("steps", 30)
    seed   = item.get("seed", -1)

    if src == "image":
        return ImageGenerationWorker(
            client=client,
            store=store,
            prompt=prompt,
            negative_prompt=neg,
            num_inference_steps=steps,
            seed=seed,
            guidance_scale=item.get("guidance_scale", 3.5),
            model=model,
        )

    if src == "animate":
        return AnimateGenerationWorker(
            client=client,
            store=store,
            reference_video_path=item.get("ref_video_path", ""),
            reference_image_path=item.get("ref_char_path", ""),
            prompt=prompt,
            num_inference_steps=steps,
            seed=seed,
            animate_mode=item.get("animate_mode", "animation"),
            model=model,
        )

    # video / mochi / skyreels → GenerationWorker
    # If a seed image path is present and the file exists, encode it as a
    # base64 data-URI so I2V models (e.g. SkyReels-V2-I2V) receive the
    # conditioning image.  The GUI does this inline; the CLI path needs to
    # replicate it here.
    seed_image_path = item.get("seed_image_path", "")
    image_b64: str | None = None
    if seed_image_path:
        _img = Path(seed_image_path)
        if _img.is_file():
            import base64 as _b64
            _mime = "image/png" if _img.suffix.lower() == ".png" else "image/jpeg"
            image_b64 = f"data:{_mime};base64,{_b64.b64encode(_img.read_bytes()).decode()}"

    return GenerationWorker(
        client=client,
        store=store,
        prompt=prompt,
        negative_prompt=neg,
        num_inference_steps=steps,
        seed=seed,
        seed_image_path=seed_image_path,
        model=model,
        num_frames=item.get("num_frames"),
        image=image_b64,
    )


# ── servers ───────────────────────────────────────────────────────────────────

def cmd_servers(args):
    """Print live status of every managed service."""
    statuses = sm.status_all(timeout=2.0)
    print(bold("Managed services"))
    for key, sdef in sm.SERVERS.items():
        alive = statuses.get(key, False)
        dot   = green("●") if alive else dim("○")
        state = green("running") if alive else dim("offline")
        print(f"  {dot}  {teal(key):<18s}  {state:<10s}  {dim(sdef.label)}")
    print()
    print(dim("Start/stop:  tt-ctl start <service>  |  tt-ctl stop <service>"))
    print(dim(f"Services:    {' | '.join(sm.SERVERS)} | all"))


def _resolve_service_arg(args) -> str:
    """Resolve service key from positional arg or --single-chip flag."""
    if getattr(args, "single_chip", False):
        return sm.ONE_CHIP_KEY
    if not args.service:
        print(red("error: specify a SERVICE or use --single-chip"))
        sys.exit(1)
    return args.service


def cmd_start(args):
    key = _resolve_service_arg(args)
    try:
        sdefs = sm._resolve(key)
    except KeyError as e:
        print(red(str(e)))
        sys.exit(1)

    for sdef in sdefs:
        print(f"Starting {teal(sdef.key)} — {sdef.label}…")

    results = sm.start(key, gui=not args.blocking)

    for sdef, result in zip(sdefs, results):
        if result.returncode == 0:
            print(green("✓") + f"  {sdef.key} started")
            if result.stdout.strip():
                for line in result.stdout.strip().splitlines()[-4:]:
                    print(dim(f"     {line}"))
        else:
            print(red("✗") + f"  {sdef.key} failed (exit {result.returncode})")
            if result.stderr.strip():
                for line in result.stderr.strip().splitlines()[-6:]:
                    print(red(f"     {line}"))


def cmd_stop(args):
    key = _resolve_service_arg(args)
    try:
        sdefs = sm._resolve(key)
    except KeyError as e:
        print(red(str(e)))
        sys.exit(1)

    for sdef in sdefs:
        print(f"Stopping {teal(sdef.key)}…")

    results = sm.stop(key)

    for sdef, result in zip(sdefs, results):
        if result.returncode == 0:
            print(green("✓") + f"  {sdef.key} stopped")
        else:
            print(red("✗") + f"  {sdef.key} stop failed (exit {result.returncode})")
            if result.stderr.strip():
                for line in result.stderr.strip().splitlines()[-4:]:
                    print(red(f"     {line}"))


def cmd_restart(args):
    key = _resolve_service_arg(args)
    try:
        sdefs = sm._resolve(key)
    except KeyError as e:
        print(red(str(e)))
        sys.exit(1)

    print(f"Restarting {teal(key)}…")
    sm.stop(key)
    print(dim("  stopped — starting…"))
    results = sm.start(key, gui=not args.blocking)

    for sdef, result in zip(sdefs, results):
        if result.returncode == 0:
            print(green("✓") + f"  {sdef.key} restarted")
        else:
            print(red("✗") + f"  {sdef.key} restart failed (exit {result.returncode})")


# ── status ────────────────────────────────────────────────────────────────────

def cmd_status(args):
    client = _make_client(args)
    store  = HistoryStore()
    queue  = store.load_queue()
    recs   = store.all_records()

    # Server
    print(bold("Server"))
    alive = client.health_check()
    if alive:
        model = client.detect_running_model()
        ready = client.model_ready()
        dot   = green("●") if ready else yellow("●")
        label = model if model else ("ready" if ready else "up, model loading")
        print(f"  {dot}  {client.base_url}  —  {label}")
    else:
        print(f"  {red('●')}  {client.base_url}  —  offline")

    # Queue
    print()
    print(bold("Queue"))
    if not queue:
        print(dim("  (empty)"))
    else:
        for i, item in enumerate(queue):
            tag   = teal(f"[{i}]")
            model = item.get("model_source", "video")
            short = item.get("prompt", "")[:80]
            if len(item.get("prompt", "")) > 80:
                short += "…"
            print(f"  {tag} {dim(model)}  {short}")

    # History summary
    videos = [r for r in recs if r.media_type == "video"]
    images = [r for r in recs if r.media_type == "image"]
    # all_records() is newest-first: recs[0] = most recent, recs[-1] = oldest
    newest = recs[0]  if recs else None
    oldest = recs[-1] if recs else None
    print()
    print(bold("History"))
    print(f"  {len(recs)} records  ({len(videos)} video, {len(images)} image)")
    if newest:
        age   = _age(newest.created_at)
        short = newest.prompt[:70] + ("…" if len(newest.prompt) > 70 else "")
        print(f"  newest: {dim(age)}  {short}")
    if oldest and oldest is not newest:
        age   = _age(oldest.created_at)
        short = oldest.prompt[:70] + ("…" if len(oldest.prompt) > 70 else "")
        print(f"  oldest: {dim(age)}  {short}")


# ── queue ─────────────────────────────────────────────────────────────────────

def cmd_queue(args):
    store = HistoryStore()
    queue = store.load_queue()

    sub = getattr(args, "queue_cmd", None)

    if sub == "add":
        item = {
            "prompt":          args.prompt,
            "negative_prompt": getattr(args, "neg", ""),
            "steps":           getattr(args, "steps", 30),
            "seed":            getattr(args, "seed", -1),
            "seed_image_path": "",
            "model_source":    getattr(args, "model", "video"),
            "guidance_scale":  5.0,
            "ref_video_path":  "",
            "ref_char_path":   "",
            "animate_mode":    "animation",
            "model_id":        "",
            "job_id_override": "",
        }
        queue.append(item)
        store.save_queue(queue)
        print(green("✓") + f"  Added to queue (position {len(queue)-1})")
        print(_wrap(args.prompt))

    elif sub == "clear":
        n = len(queue)
        store.save_queue([])
        print(green("✓") + f"  Cleared {n} item(s) from queue")

    elif sub == "run":
        cmd_queue_run(args)

    else:
        # default: list
        if not queue:
            print(dim("Queue is empty."))
            return
        print(bold(f"Queue  ({len(queue)} item{'s' if len(queue)!=1 else ''})"))
        for i, item in enumerate(queue):
            model = item.get("model_source", "video")
            steps = item.get("steps", "?")
            seed  = item.get("seed", -1)
            print(f"\n  {teal(f'[{i}]')} {dim(model)}  steps={steps}  seed={seed}")
            print(_wrap(item.get("prompt", "(no prompt)")))


# ── queue run ─────────────────────────────────────────────────────────────────

def cmd_queue_run(args):
    """
    Drain and execute the persistent queue (blocking).

    Items are removed from queue.json as they complete (success or failure).
    Server-offline items are skipped and left in the queue for later.
    On KeyboardInterrupt, the in-flight item is removed (it was submitted to the
    server and can be recovered with tt-ctl recover); remaining items stay.
    """
    store  = HistoryStore()
    client = _make_client(args)
    dry_run = getattr(args, "dry_run", False)

    queue = store.load_queue()
    if not queue:
        print(dim("Queue is empty."))
        return

    total = len(queue)
    print(bold(f"Queue: {total} item{'s' if total != 1 else ''}"))

    if dry_run:
        for i, item in enumerate(queue):
            src   = item.get("model_source", "video")
            steps = item.get("steps", "?")
            seed  = item.get("seed", -1)
            short = item.get("prompt", "")[:72]
            print(f"  {teal(f'[{i}]')} {dim(src)}  steps={steps}  seed={seed}")
            print(_wrap(short))
        return

    n_done = n_failed = n_skipped = 0
    # skipped_items: server-offline items preserved in queue after the run.
    skipped_items: list = []

    for idx, item in enumerate(queue):
        src   = item.get("model_source", "video")
        short = item.get("prompt", "")[:60]
        print(f"\n  {teal(f'[{idx+1}/{total}]')} {dim(src)}  {short}")

        if not client.health_check():
            print(yellow("    ○ Server offline — skipping (item stays in queue)"))
            n_skipped += 1
            skipped_items.append(item)
            continue

        if src not in _MODEL_FOR_SRC:
            print(yellow(f"    ! Unknown model_source '{src}' — skipping (item stays in queue)"))
            n_skipped += 1
            skipped_items.append(item)
            continue

        worker = _make_worker_for_item(client, store, item)

        done   = threading.Event()
        result: dict = {}

        def on_progress(msg):
            print(dim(f"    {msg}"))

        def on_finished(record):
            result["record"] = record
            done.set()

        def on_error(msg):
            result["error"] = msg
            done.set()

        t = threading.Thread(
            target=worker.run_with_callbacks,
            kwargs=dict(on_progress=on_progress, on_finished=on_finished,
                        on_error=on_error),
            daemon=True,
        )
        t.start()

        try:
            done.wait()
        except KeyboardInterrupt:
            worker.cancel()
            # Current item was submitted to server — remove it from queue.
            # Remaining unprocessed items + any previously skipped items are kept.
            remaining = skipped_items + list(queue[idx + 1:])
            store.save_queue(remaining)
            jid = getattr(worker, "_current_job_id", None) or "unknown"
            jid_short = jid[:8] if jid != "unknown" else jid
            print(yellow(f"\n  Cancelled. Job may still be running on server ({jid_short}…)"))
            print(dim("  Run: tt-ctl recover   to retrieve any completed result"))
            sys.exit(1)

        if "error" in result:
            print(red(f"    ✗ {result['error']}"))
            n_failed += 1
        else:
            record = result["record"]
            print(green("    ✓") + f"  {record.media_file_path}")
            try:
                dur_str = f"{record.duration_s:.0f}s"
            except (TypeError, ValueError):
                dur_str = str(record.duration_s)
            print(dim(f"       {dur_str}  |  {record.model}"))
            n_done += 1

        # Remove this item from queue.json after each completion (success or fail).
        # skipped_items + everything after current index = what remains.
        store.save_queue(skipped_items + list(queue[idx + 1:]))

    # Restore any server-offline items at their original positions (front of queue).
    store.save_queue(skipped_items)

    print()
    parts = []
    if n_done:    parts.append(green(f"{n_done} done"))
    if n_failed:  parts.append(red(f"{n_failed} failed"))
    if n_skipped: parts.append(yellow(f"{n_skipped} skipped"))
    print("  " + "  ".join(parts) if parts else dim("  Nothing ran."))


# ── generate ──────────────────────────────────────────────────────────────────

def cmd_generate(args):
    """
    Generate N prompts (optionally guided by a theme) and run them through the queue.

    Prompts are written to queue.json first, then cmd_queue_run drains the queue
    unless --queue-only is set.  With an explicit --seed S and --count N, seeds
    are assigned as S, S+1, S+2, … so each generation gets a unique seed while
    still being reproducible when re-run with the same starting seed.
    """
    guide      = getattr(args, "guide", None)
    count      = getattr(args, "count", 1)
    ptype      = getattr(args, "type", "video")
    mode       = getattr(args, "mode", "algo")
    enhance    = not getattr(args, "no_enhance", False)
    steps      = getattr(args, "steps", 30)
    seed       = getattr(args, "seed", -1)
    queue_only = getattr(args, "queue_only", False)

    store          = HistoryStore()
    existing_queue = store.load_queue()

    # Warn once at the start if LLM polish / guided generation was requested but
    # the Qwen server is not available, so the user knows the fallback is active.
    if enhance:
        if not gp._llm_available():
            print(yellow("  Qwen server offline — algo fallback active"))

    print(bold(f"Generating {count} prompt{'s' if count != 1 else ''}…"))

    new_items = []
    for i in range(count):
        # Explicit seed: each item gets seed+i so runs are unique yet reproducible.
        # Random seed (-1): all items keep -1 so the server picks a random seed
        # independently for each generation.
        item_seed = (seed + i) if seed >= 0 else -1

        if guide:
            result = gp.guided_generate(guide, ptype, enhance=enhance)
        else:
            result = gp.generate(prompt_type=ptype, mode=mode, enhance=enhance)

        prompt = result["prompt"]
        source = result["source"]
        short  = prompt[:70] + ("…" if len(prompt) > 70 else "")
        print(f"  {dim(f'[{i+1}/{count}]')} {dim(source)}  {short}")

        new_items.append({
            "prompt":          prompt,
            "negative_prompt": "",
            "steps":           steps,
            "seed":            item_seed,
            "seed_image_path": "",
            "model_source":    ptype,
            "guidance_scale":  5.0,
            "ref_video_path":  "",
            "ref_char_path":   "",
            "animate_mode":    "animation",
            "model_id":        "",
            "job_id_override": "",
        })

    store.save_queue(existing_queue + new_items)
    print(green("✓") + f"  {count} item{'s' if count != 1 else ''} added to queue")

    if queue_only:
        print(dim("  Run: tt-ctl queue run   to start generation"))
        return

    cmd_queue_run(args)


# ── history ───────────────────────────────────────────────────────────────────

def cmd_history(args):
    store = HistoryStore()
    recs  = store.all_records()
    n     = getattr(args, "n", 10)
    shown = recs[:n]          # all_records() is already newest-first

    if not shown:
        print(dim("No history yet."))
        return

    print(bold(f"Last {len(shown)} generation(s)"))
    for r in shown:
        age    = _age(r.created_at)
        exists = green("✓") if r.media_exists else red("✗")
        model  = dim(r.model or r.media_type)
        gen    = dim(f"{r.duration_s:.0f}s gen") if r.duration_s else ""
        print(f"\n  {exists} {dim(age)}  {model}  {gen}")
        print(_wrap(r.prompt))
        if r.media_exists:
            print(dim(f"     {r.media_file_path}"))
        else:
            print(red(f"     FILE MISSING: {r.media_file_path}"))


# ── server ────────────────────────────────────────────────────────────────────

def cmd_server(args):
    client = _make_client(args)

    alive = client.health_check()
    if not alive:
        print(f"{red('●')}  {client.base_url}  —  offline")
        return

    model = client.detect_running_model()
    ready = client.model_ready()
    dot   = green("●") if ready else yellow("●")
    print(f"{dot}  {client.base_url}")
    print(f"   model : {model or dim('(not detected)')}")
    print(f"   ready : {'yes' if ready else 'loading…'}")

    try:
        jobs = client.list_jobs()
        print(f"   jobs  : {len(jobs)} on server")
        for j in jobs[:5]:
            jid    = j.get("id", "?")[:12]
            status = j.get("status", "?")
            prompt = (j.get("prompt") or j.get("request", {}).get("prompt") or "")[:60]
            col    = green if status in ("completed", "succeeded") else (
                      red   if status in ("failed", "error")       else yellow)
            print(f"           {col(status):12s}  {dim(jid)}  {prompt}")
        if len(jobs) > 5:
            print(dim(f"           … and {len(jobs)-5} more"))
    except Exception as exc:
        print(dim(f"   (could not list jobs: {exc})"))


# ── logs ─────────────────────────────────────────────────────────────────────

# Maps service key → glob pattern under workflow_logs/docker_server/
_LOG_GLOBS: dict[str, str] = {
    "wan2.2":         "media_*_Wan2.2-T2V-A14B-Diffusers_p300x2_server.log",
    "mochi":          "media_*_mochi-1-preview_p300x2_server.log",
    "flux":           "media_*_FLUX.1-dev_p300x2_server.log",
    "animate":        "media_*_Wan2.2-Animate-14B-Diffusers_p300x2_server.log",
    "skyreels":       "media_*_SkyReels-V2-I2V-14B-540P_*_server.log",
    "z-image-turbo":  "media_*_Z-Image-Turbo_p150x4_server.log",
    "motif":          "media_*_Motif-Image-6B-Preview_p300x2_server.log",
}
_LOG_DIR = _HERE / "workflow_logs" / "docker_server"
_PROMPT_LOG = Path("/tmp/tt_prompt_gen.log")


def _latest_log_file(service_key: str) -> Path | None:
    """Return the most-recently-modified log file for a service, or None."""
    pattern = _LOG_GLOBS.get(service_key)
    if not pattern:
        return None
    matches = sorted(
        _glob.glob(str(_LOG_DIR / pattern)),
        key=lambda p: Path(p).stat().st_mtime,
        reverse=True,
    )
    return Path(matches[0]) if matches else None


def _running_container_id() -> str | None:
    """Return the ID of the first running inference server container, if any."""
    try:
        result = subprocess.run(
            ["docker", "ps", "--filter", "name=tt-inference-server",
             "--format", "{{.ID}}"],
            capture_output=True, text=True, timeout=5,
        )
        ids = result.stdout.strip().splitlines()
        if ids:
            return ids[0]
        # Fallback: any container from the known image
        result2 = subprocess.run(
            ["docker", "ps", "--filter",
             "ancestor=ghcr.io/tenstorrent/tt-media-inference-server",
             "--format", "{{.ID}}"],
            capture_output=True, text=True, timeout=5,
        )
        ids2 = result2.stdout.strip().splitlines()
        return ids2[0] if ids2 else None
    except Exception:
        return None


def cmd_logs(args):
    """
    Tail the log for a managed service.

    Source priority:
      - prompt-server        → /tmp/tt_prompt_gen.log  (static file)
      - docker               → docker logs of the running container (live stream)
      - wan2.2 / mochi / flux / animate / skyreels
          → newest file in workflow_logs/docker_server/ matching the service
            pattern, falling back to docker logs if no file found
      - (no service given)   → docker logs of the running container
    """
    service  = getattr(args, "service", None)
    follow   = getattr(args, "follow", False)
    n_lines  = getattr(args, "lines", 50)

    # ── prompt-server: static file ────────────────────────────────────────────
    if service == "prompt-server":
        if not _PROMPT_LOG.exists():
            print(yellow(f"Prompt server log not found: {_PROMPT_LOG}"))
            print(dim("  Start the prompt server first:  tt-ctl start prompt-server"))
            return
        print(dim(f"Log: {_PROMPT_LOG}"))
        _tail_file(_PROMPT_LOG, n_lines, follow)
        return

    # ── docker: stream container logs directly ────────────────────────────────
    if service == "docker" or service is None:
        cid = _running_container_id()
        if not cid:
            print(yellow("No running inference server container found."))
            print(dim("  Start a server first:  tt-ctl start skyreels"))
            return
        print(dim(f"Container: {cid}"))
        _docker_logs(cid, n_lines, follow)
        return

    # ── inference services: file log, fall back to docker ────────────────────
    log_file = _latest_log_file(service)
    if log_file:
        print(dim(f"Log: {log_file}"))
        _tail_file(log_file, n_lines, follow)
    else:
        # No log file yet — try the running container
        cid = _running_container_id()
        if not cid:
            print(yellow(f"No log file found for '{service}' and no container running."))
            return
        print(dim(f"No log file found for '{service}' — streaming container logs ({cid})"))
        _docker_logs(cid, n_lines, follow)


def _tail_file(path: Path, n_lines: int, follow: bool) -> None:
    """Print last n_lines of a file, optionally following."""
    cmd = ["tail", f"-n{n_lines}"]
    if follow:
        cmd.append("-f")
    cmd.append(str(path))
    try:
        subprocess.run(cmd)
    except KeyboardInterrupt:
        pass


def _docker_logs(container_id: str, n_lines: int, follow: bool) -> None:
    """Stream docker logs for a container."""
    cmd = ["docker", "logs", f"--tail={n_lines}"]
    if follow:
        cmd.append("--follow")
    cmd.append(container_id)
    try:
        subprocess.run(cmd)
    except KeyboardInterrupt:
        pass


# ── recover ───────────────────────────────────────────────────────────────────

def cmd_recover(args):
    """List server jobs whose IDs don't appear in local history.

    Does not modify anything — prints the job IDs and prompts so you can
    decide what to do.  To actually recover, the GUI's Recover button uses
    these same IDs via the AttachRecoveryJob flow.
    """
    client = _make_client(args)
    store  = HistoryStore()

    alive = client.health_check()
    if not alive:
        print(red("Server offline — cannot scan jobs."))
        sys.exit(1)

    known_ids = {r.id for r in store.all_records()}

    try:
        jobs = client.list_jobs()
    except Exception as exc:
        print(red(f"Failed to list jobs: {exc}"))
        sys.exit(1)

    unknown = [j for j in jobs if j.get("id") not in known_ids]

    if not unknown:
        print(green("✓") + "  All server jobs are in local history.")
        return

    print(bold(f"{len(unknown)} unknown job(s) on server:"))
    for j in unknown:
        jid    = j.get("id", "?")
        status = j.get("status", "?")
        prompt = (j.get("prompt") or j.get("request", {}).get("prompt") or "")[:80]
        col    = green if status in ("completed", "succeeded") else (
                  red   if status in ("failed", "error")       else yellow)
        print(f"\n  {col(status):12s}  {dim(jid)}")
        if prompt:
            print(_wrap(prompt))

    print()
    print(dim("To recover: open the GUI and click '⟳ Recover Jobs', or requeue "
              "the prompt with:  tt-ctl queue add \"<prompt>\""))


# ── run ───────────────────────────────────────────────────────────────────────

def cmd_run(args):
    """Run one generation synchronously, writing to the shared history."""
    client = _make_client(args)

    alive = client.health_check()
    if not alive:
        print(red("Server offline."))
        sys.exit(1)

    if not client.model_ready():
        print(yellow("Server is up but model is still loading — try again shortly."))
        sys.exit(1)

    store  = HistoryStore()
    prompt = args.prompt
    steps  = getattr(args, "steps", 30)
    seed   = getattr(args, "seed",  -1)
    neg    = getattr(args, "neg",   "")
    model_src = getattr(args, "model", "video")

    print(bold("Submitting generation"))
    print(_wrap(prompt))
    print(dim(f"  steps={steps}  seed={seed}  model={model_src}"))
    print()

    done   = threading.Event()
    result = {}

    def on_progress(msg: str):
        print(dim(f"  {msg}"))

    def on_finished(record):
        result["record"] = record
        done.set()

    def on_error(msg: str):
        result["error"] = msg
        done.set()

    worker = GenerationWorker(
        client=client,
        store=store,
        prompt=prompt,
        negative_prompt=neg,
        num_inference_steps=steps,
        seed=seed,
        model=model_src,
    )
    t = threading.Thread(
        target=worker.run_with_callbacks,
        kwargs=dict(
            on_progress=on_progress,
            on_finished=on_finished,
            on_error=on_error,
        ),
        daemon=True,
    )
    t.start()

    try:
        done.wait()
    except KeyboardInterrupt:
        worker.cancel()
        print(yellow("\nCancelled."))
        sys.exit(1)

    if "error" in result:
        print(red(f"\nError: {result['error']}"))
        sys.exit(1)

    record = result["record"]
    # worker already called store.append(record) internally
    print(green("\n✓ Done"))
    print(f"  {record.media_file_path}")
    print(dim(f"  gen time: {record.duration_s:.0f}s  |  model: {record.model}"))


# ── serve-inventory ───────────────────────────────────────────────────────────

def cmd_serve_inventory(args):
    """Serve the local video library over HTTP for a remote GUI to browse.

    Start this command on the machine that has the generated videos.  Then
    run the GUI on another machine with::

        ./tt-gen --server http://this-machine:8000

    The GUI derives the inventory URL as http://this-machine:8002 and fetches
    records automatically at startup.
    """
    import importlib, sys as _sys, logging as _logging
    _logging.basicConfig(level=_logging.INFO, format="%(message)s")

    # Import inventory_server from app/
    _sys.path.insert(0, str(_HERE / "app"))
    inv = importlib.import_module("inventory_server")

    port = getattr(args, "port", inv.DEFAULT_PORT)
    host = getattr(args, "host", "0.0.0.0")
    inv.serve(port=port, host=host)


# ── arg parsing ───────────────────────────────────────────────────────────────

def _build_parser() -> argparse.ArgumentParser:
    p = argparse.ArgumentParser(
        prog="tt-ctl",
        description="CLI for tt-local-generator — works without the GUI.",
        formatter_class=argparse.RawDescriptionHelpFormatter,
    )
    p.add_argument("--version", action="version", version=f"tt-ctl {_VERSION}")
    p.add_argument("--server", default="http://localhost:8000",
                   help="Inference server URL (default: http://localhost:8000)")

    sub = p.add_subparsers(dest="cmd", metavar="COMMAND")

    # status
    sub.add_parser("status", help="Server health, queue depth, history summary")

    # servers
    sub.add_parser("servers", help="Live status of every managed service")

    # start / stop / restart
    _service_choices = sorted(sm.SERVERS.keys()) + [sm.ALL_KEY]
    _service_help = " | ".join(_service_choices)

    for _verb in ("start", "stop", "restart"):
        _sp = sub.add_parser(_verb, help=f"{_verb.capitalize()} a managed service")
        _sp.add_argument("service", nargs="?", choices=_service_choices, metavar="SERVICE",
                         help=_service_help)
        _sp.add_argument("--single-chip", action="store_true", dest="single_chip",
                         help="Start artgen + prompt-server (single Blackhole card or CPU-only)")
        if _verb != "stop":
            _sp.add_argument("--blocking", action="store_true",
                             help="Wait for the script to exit (default: --gui mode)")

    # queue
    q = sub.add_parser("queue", help="Manage the persistent queue")
    qsub = q.add_subparsers(dest="queue_cmd", metavar="SUBCOMMAND")
    qa = qsub.add_parser("add", help="Append a prompt to the queue")
    qa.add_argument("prompt")
    qa.add_argument("--model", default="video",
                    choices=["video", "image", "animate", "skyreels", "mochi"],
                    help="Model type (default: video)")
    qa.add_argument("--steps", type=int, default=30)
    qa.add_argument("--seed",  type=int, default=-1)
    qa.add_argument("--neg",   default="", metavar="NEG_PROMPT")
    qsub.add_parser("clear", help="Empty the queue")
    qr = qsub.add_parser("run", help="Drain and execute the queue (blocking)")
    qr.add_argument("--dry-run", action="store_true",
                    help="Print what would run without executing")

    # history
    h = sub.add_parser("history", help="Show recent completed generations")
    h.add_argument("n", nargs="?", type=int, default=10,
                   metavar="N", help="How many to show (default 10)")

    # server
    sub.add_parser("server", help="Server detail and job list")

    # logs
    _log_service_choices = sorted(_LOG_GLOBS.keys()) + ["prompt-server", "docker"]
    lg = sub.add_parser("logs", help="Tail logs for a managed service or container")
    lg.add_argument(
        "service", nargs="?", default=None,
        choices=_log_service_choices,
        metavar="SERVICE",
        help=f"Service to show logs for ({' | '.join(_log_service_choices)}); "
             "default: running container",
    )
    lg.add_argument("-f", "--follow", action="store_true",
                    help="Follow / tail the log in real time (Ctrl-C to stop)")
    lg.add_argument("-n", "--lines", type=int, default=50, metavar="N",
                    help="Number of lines to show (default 50)")

    # recover
    sub.add_parser("recover", help="List server jobs not in local history")

    # artgen
    _build_artgen_parser(sub)

    # plugin management
    plugin_p = sub.add_parser("plugin", help="Plugin management")
    plugin_sub = plugin_p.add_subparsers(dest="plugin_cmd")
    plugin_sub.add_parser("list", help="List all loaded plugins")

    # Workflow runner
    wf_p = sub.add_parser("workflow", help="Run a multi-step generative workflow")
    wf_sub = wf_p.add_subparsers(dest="workflow_cmd")
    wf_run = wf_sub.add_parser("run", help="Execute a workflow JSON spec")
    wf_run.add_argument("spec", help="Path to workflow JSON (e.g. docs/examples/workflows/1964-worlds-fair.json)")
    wf_run.add_argument("--dry-run", action="store_true", help="Print steps without running inference")
    wf_sub.add_parser("list", help="List available workflow examples")

    # MCP config
    sub.add_parser("mcp-config", help="Emit Claude Code MCP config JSON for tt-local-gen")

    # serve-inventory
    inv = sub.add_parser(
        "serve-inventory",
        help="Share the local video library over HTTP for a remote GUI to browse",
    )
    inv.add_argument(
        "--port", type=int, default=8002,
        help="TCP port to listen on (default 8002)",
    )
    inv.add_argument(
        "--host", default="0.0.0.0",
        help="Interface to bind (default 0.0.0.0 = all interfaces)",
    )

    # run
    r = sub.add_parser("run", help="Run one generation now (blocking)")
    r.add_argument("prompt")
    r.add_argument("--steps", type=int, default=30)
    r.add_argument("--seed",  type=int, default=-1)
    r.add_argument("--neg",   default="", metavar="NEG_PROMPT")
    r.add_argument("--model", default="video",
                   choices=["video", "image", "animate", "skyreels", "mochi"])

    # generate
    gen = sub.add_parser(
        "generate",
        help="Generate prompts (optionally guided) and run them",
    )
    gen.add_argument(
        "guide", nargs="?", default=None, metavar="GUIDE",
        help="Optional guiding theme — Qwen generates around it if available",
    )
    gen.add_argument("--count", type=int, default=1, metavar="N",
                     help="Number of prompts to generate (default: 1)")
    gen.add_argument(
        "--type", default="video",
        choices=["video", "image", "skyreels", "animate"],
        metavar="TYPE",
        help="Model type: video|image|skyreels|animate (default: video)",
    )
    gen.add_argument(
        "--mode", default="algo", choices=["algo", "markov"],
        help="Base generation mode (default: algo; ignored when GUIDE + Qwen up)",
    )
    gen.add_argument(
        "--no-enhance", action="store_true",
        help="Skip Qwen polish / guided generation even if server is up",
    )
    gen.add_argument("--steps", type=int, default=30,
                     help="Inference steps per generation (default: 30)")
    gen.add_argument("--seed", type=int, default=-1,
                     help="Starting seed (-1 = random). With --count N, seeds are S, S+1, …")
    gen.add_argument("--queue-only", action="store_true",
                     help="Stage prompts in queue.json but do not run them")
    gen.add_argument("--server", default="http://localhost:8000",
                     help="Inference server URL (default: http://localhost:8000)")
    gen.add_argument("--dry-run", action="store_true",
                     help="Preview queued items without executing")

    return p


def main():
    parser = _build_parser()
    args   = parser.parse_args()

    if args.cmd is None:
        parser.print_help()
        return

    dispatch = {
        "status":   cmd_status,
        "servers":  cmd_servers,
        "start":    cmd_start,
        "stop":     cmd_stop,
        "restart":  cmd_restart,
        "queue":    cmd_queue,
        "history":  cmd_history,
        "server":   cmd_server,
        "logs":     cmd_logs,
        "recover":          cmd_recover,
        "run":              cmd_run,
        "generate":         cmd_generate,
        "serve-inventory":  cmd_serve_inventory,
        "artgen":           cmd_artgen,
        "plugin":           cmd_plugin,
        "mcp-config":       cmd_mcp_config,
        "workflow":         cmd_workflow,
    }
    dispatch[args.cmd](args)


if __name__ == "__main__":
    main()
