#!/usr/bin/env python3
"""Generate a LoRA trigger word dictionary HTML page from .safetensors metadata.

Usage:
    lora-trigger-sheet                                    # default: ~/mnt/nimo.loc/ComfyUI/models/loras
    lora-trigger-sheet /path/to/loras                     # specify directory
    lora-trigger-sheet .                                  # current directory
    lora-trigger-sheet /path/to/loras /path/to/output     # specify both input and output
"""

import json
import os
import re
import struct
import sys
from pathlib import Path

import hashlib
import time
from urllib.error import HTTPError, URLError
from urllib.parse import quote
from urllib.request import urlopen, Request


def read_safetensors_meta(filepath):
    try:
        with open(filepath, "rb") as f:
            header_len = struct.unpack("<Q", f.read(8))[0]
            raw = f.read(header_len)
            parsed = json.loads(raw)
    except Exception:
        return None

    # Metadata is inside "__metadata__" key
    if "__metadata__" in parsed:
        return parsed["__metadata__"]
    return parsed


def get_tool(meta):
    if meta is None:
        return "error"

    if meta.get("ss_sd_scripts_commit_hash"):
        return "sd-scripts"

    net = (meta.get("ss_training_network") or meta.get("ss_network_module") or "").lower()
    if "kohya" in net:
        return "kohya"
    if "sd-scripts" in net or "sd_scripts" in net:
        return "sd-scripts"

    if meta.get("ss_learning_rate"):
        return "sd-scripts"

    if net:
        return net

    return "unknown"


def get_base_model(meta):
    if meta is None:
        return "error"

    base = (meta.get("ss_base_model_version")
            or meta.get("ss_base_model")
            or meta.get("modelspec.architecture", "")
            .replace("/lora", "")
            or "")
    if base:
        return base

    return "unknown"


def get_tags(meta):
    if meta is None:
        return None, 0

    raw = meta.get("ss_tag_frequency")
    if not raw:
        return None, 0

    # ss_tag_frequency is often a JSON string; parse if needed
    if isinstance(raw, str):
        try:
            raw = json.loads(raw)
        except (json.JSONDecodeError, TypeError):
            return None, 0

    tags = {}
    total = 0

    if isinstance(raw, dict):
        for category, tag_dict in raw.items():
            if isinstance(tag_dict, dict):
                for tag, count in tag_dict.items():
                    total += count
                    tags[tag] = tags.get(tag, 0) + count
            elif isinstance(tag_dict, (int, float)):
                total += tag_dict

    if not tags:
        if all(isinstance(v, (int, float)) for v in raw.values()):
            tags = raw
            total = sum(raw.values())
        else:
            return None, 0

    return tags, total


def shorten_base_model(name):
    name = re.sub(r"[-_][Vv]\d+(?:[._-]\d+)?$", "", name)
    name = re.sub(r"[-_][Vv]\d+[a-z]?\d*$", "", name)
    replacements = {
        "stable-diffusion-xl-v1-base": "sdxl_base_v1-0",
        "stable diffusion xl base v1.0": "sdxl_base_v1-0",
        "sdxl base v1.0": "sdxl_base_v1-0",
        "sdxl base v0.9": "sdxl_base_v0-9",
        "flux.1": "flux1",
        "flux": "flux1",
    }
    for old, new in replacements.items():
        if name.lower() == old.lower() or name.lower().startswith(old.lower()):
            return new
    return name


def html_id():
    import hashlib, time
    return hashlib.md5(str(time.monotonic_ns()).encode()).hexdigest()[:8]


# --- Source URL lookup cache ---

CACHE_DIR = os.path.expanduser("~/.cache")
CACHE_FILE = os.path.join(CACHE_DIR, "lora-sources.json")


def load_source_cache():
    if os.path.exists(CACHE_FILE):
        try:
            with open(CACHE_FILE) as f:
                return json.load(f)
        except (json.JSONDecodeError, IOError):
            return {}
    return {}


def save_source_cache(cache):
    os.makedirs(CACHE_DIR, exist_ok=True)
    with open(CACHE_FILE, "w") as f:
        json.dump(cache, f, indent=2)


def file_blake2b_hash(filepath):
    h = hashlib.blake2b(digest_size=32)
    with open(filepath, "rb") as f:
        while True:
            chunk = f.read(65536)
            if not chunk:
                break
            h.update(chunk)
    return h.hexdigest()


def _fetch_json(url, timeout=10):
    req = Request(url, headers={"User-Agent": "lora-trigger-sheet/1.0"})
    with urlopen(req, timeout=timeout) as resp:
        return json.loads(resp.read().decode())


def lookup_source(filepath, filename, cache):
    now = int(time.time())
    cached = cache.get(filename)
    if cached:
        url = cached.get("url")
        label = cached.get("label")
        last = cached.get("last_checked", 0)
        if (now - last) < 86400 * 30:
            return (url, label) if url else (None, None)

    # CivitAI model-version lookup by BLAKE2b file hash
    try:
        h = file_blake2b_hash(filepath)
        data = _fetch_json(
            f"https://civitai.com/api/v1/model-versions/by-hash/{h}"
        )
        if isinstance(data, dict):
            model_id = data.get("modelId")
            version_id = data.get("id")
            if model_id:
                url = f"https://civitai.com/models/{model_id}"
                if version_id:
                    url += f"?modelVersionId={version_id}"
                cache[filename] = {
                    "url": url, "label": "CivitAI",
                    "model_name": data.get("name", ""),
                    "last_checked": now,
                }
                return url, "CivitAI"
    except (HTTPError, URLError, json.JSONDecodeError, OSError):
        pass

    # Fallback: HuggingFace model search by filename
    try:
        q = (
            filename.replace(".safetensors", "")
            .replace("_", " ")
            .replace("-", " ")
        )
        data = _fetch_json(
            "https://huggingface.co/api/models"
            f"?search={quote(q)}&sort=downloads&direction=-1&limit=3"
        )
        if isinstance(data, list) and len(data) > 0:
            model_id = data[0].get("modelId", "")
            if model_id:
                url = f"https://huggingface.co/{model_id}"
                cache[filename] = {
                    "url": url, "label": "HuggingFace",
                    "model_name": model_id, "last_checked": now,
                }
                return url, "HuggingFace"
    except (HTTPError, URLError, json.JSONDecodeError, OSError):
        pass

    cache[filename] = {"url": None, "label": None, "last_checked": now}
    return None, None


def generate_html(loras, sources=None):
    if sources is None:
        sources = {}
    rows = []
    no_trigger_files = []

    for fname, meta in loras:
        tool = get_tool(meta)
        base = shorten_base_model(get_base_model(meta))
        tags, total = get_tags(meta)

        fname_no_ext = re.sub(r"\.safetensors$", "", fname)

        if tags:
            sorted_tags = sorted(tags.items(), key=lambda x: -x[1])
            top8 = sorted_tags[:8]
            rest = sorted_tags[8:]
            has_rest = len(rest) > 0

            top8_html = ", ".join(
                f"<code>{_e(t)}</code> <span class=\"c\">({c}x)</span>"
                for t, c in top8
            )
            rest_html = ""
            more_id = ""
            if has_rest:
                more_id = "m_" + html_id()
                more_count = sum(c for _, c in rest)
                rest_html = " <span class=\"ml\" onclick=\"t('" + more_id + "',this)\" data-n=\"" + str(more_count) + "\">" + str(more_count) + " more&hellip;</span>"
                rest_list = ", ".join(
                    f"<code>{_e(t)}</code> <span class=\"c\">({c}x)</span>"
                    for t, c in rest
                )
                rest_html += "\n          <span class=\"ht\" id=\"" + more_id + "\">" + rest_list + "</span>"

            triggers = top8_html + rest_html
        else:
            no_trigger_files.append((fname_no_ext, tool, base))
            triggers = "<span class=\"n\">(none)</span>"

        src_url, src_label = sources.get(fname, (None, None))
        rows.append((fname_no_ext, tool, base, triggers, src_url or "", src_label or ""))

    num_files = len(rows)

    table_rows = ""
    for fname, tool, base, triggers, src_url, src_label in rows:
        if src_url:
            src_cell = f'<a href="{_e(src_url)}" target="_blank" rel="noopener">{_e(src_label)}</a>'
        else:
            src_cell = '<span class="n">—</span>'
        table_rows += (
            f"    <tr>\n"
            f"      <td class=\"fn\">{_e(fname)}</td>\n"
            f"      <td class=\"tl\">{_e(tool)}</td>\n"
            f"      <td class=\"bm\">{_e(base)}</td>\n"
            f"      <td class=\"tr\">{triggers}</td>\n"
            f"      <td class=\"sr\">{src_cell}</td>\n"
            f"    </tr>\n"
        )

    html = f"""<!DOCTYPE html>
<html lang="en">
<head><meta charset="UTF-8"><meta name="viewport" content="width=device-width,initial-scale=1.0">
<title>LoRA Info Sheet</title>
<link rel="icon" type="image/png" href="@img/nude.png">
<style>
*,*::before,*::after{{box-sizing:border-box;margin:0;padding:0}}
body{{font-family:-apple-system,BlinkMacSystemFont,'Segoe UI',Roboto,sans-serif;background:#0d1117;color:#c9d1d9;padding:24px;line-height:1.5;font-size:14px}}
h1{{font-size:1.5rem;margin-bottom:4px;color:#f0f6fc}}
.sub{{color:#8b949e;margin-bottom:16px;font-size:0.85rem}}
.bar{{display:flex;gap:12px;margin-bottom:12px;flex-wrap:wrap;align-items:center}}
#q{{flex:1;min-width:200px;padding:8px 12px;border:1px solid #30363d;border-radius:6px;background:#161b22;color:#c9d1d9;font-size:0.85rem;outline:none}}
#q:focus{{border-color:#58a6ff}}
#cnt{{color:#8b949e;font-size:0.85rem}}
table{{width:100%;border-collapse:collapse;font-size:0.82rem}}
th{{text-align:left;padding:9px 10px;border-bottom:2px solid #30363d;color:#f0f6fc;font-weight:600;cursor:pointer;user-select:none;white-space:nowrap;background:#161b22;position:sticky;top:0;z-index:1}}
th:hover{{color:#58a6ff}}
th .ar{{margin-left:3px;color:#58a6ff}}
td{{padding:7px 10px;border-bottom:1px solid #21262d;vertical-align:top;word-break:break-word}}
tr:hover td{{background:#161b22}}
.fn{{color:#58a6ff;font-weight:500;white-space:nowrap}}
.tl,.bm{{color:#8b949e}}
.tr code{{background:#1f2937;padding:1px 5px;border-radius:3px;font-size:0.78rem;color:#f0c674;white-space:nowrap}}
.c{{color:#8b949e;font-size:0.72rem;margin-right:3px}}
.n{{color:#484f58;font-style:italic}}
.ml{{color:#58a6ff;cursor:pointer;font-size:0.78rem;text-decoration:none;border-bottom:1px dotted #58a6ff}}
.ml:hover{{border-bottom:1px solid #58a6ff}}
.ht{{display:none}}
mark.hl{{background:#264f78;color:#f0f6fc;border-radius:2px;padding:0 2px}}
.sr a{{color:#58a6ff;text-decoration:none;font-size:0.78rem}}
.c0{{width:24%}}.c1{{width:8%}}.c2{{width:12%}}.c3{{width:48%}}.c4{{width:8%}}
@media(max-width:768px){{.c1,.c2,.c4,td:nth-child(2),td:nth-child(3),td:nth-child(5){{display:none}}.c0{{width:30%}}.c3{{width:70%}}}}
.foot{{margin-top:20px;color:#8b949e}}
.foot h2{{font-size:1rem;color:#f0f6fc;margin-bottom:6px}}
.foot ul{{list-style:none;columns:3;font-size:0.82rem}}
.foot li{{padding:1px 0}}
</style></head>
<body>
<h1>LoRA Info Sheet</h1>
<p class="sub">{num_files} LoRA files &middot; <span id="vc">{num_files}</span> visible</p>
<div class="bar"><input type="text" id="q" placeholder="Search filename, trigger, base model..." autofocus> <span id="cnt">{num_files} / {num_files}</span></div>
<table><thead><tr>
<th class="c0" onclick="s(0)">File <span class="ar">&#9650;</span></th>
<th class="c1" onclick="s(1)">Tool <span class="ar"></span></th>
<th class="c2" onclick="s(2)">Base <span class="ar"></span></th>
<th class="c3" onclick="s(3)">Trigger(s) <span class="ar"></span></th>
<th class="c4" onclick="s(4)">Source <span class="ar"></span></th>
</tr></thead><tbody id="b">
{table_rows}</tbody></table>

<script>
let col = 0, desc = true;
const b = document.getElementById('b');
const q = document.getElementById('q');
const cnt = document.getElementById('cnt');
const vc = document.getElementById('vc');

function s(i){{
    if (col == i) desc = !desc; else {{ col = i; desc = i == 0; }}
    let rows = [...b.children];
    rows.sort((a,c)=>{{
        let av = a.children[col].textContent.trim().toLowerCase();
        let cv = c.children[col].textContent.trim().toLowerCase();
        if (col == 0 || col == 3) {{
            let am = parseInt(av.match(/\\d+/)?.[0]||'0');
            let cm = parseInt(cv.match(/\\d+/)?.[0]||'0');
            if (!isNaN(am) && !isNaN(cm)) return desc ? cm-am : am-cm;
        }}
        return desc ? av.localeCompare(cv) : cv.localeCompare(av);
    }});
    rows.forEach(r => b.appendChild(r));
    let arrows = document.querySelectorAll('.ar');
    arrows.forEach(a => a.innerHTML = '');
    document.querySelectorAll('th')[col+1].querySelector('.ar').innerHTML = desc ? '&#9650;' : '&#9660;';
}}

function t(id,el){{
    let sp = document.getElementById(id);
    if (sp.classList.contains('ht')){{ sp.classList.remove('ht'); el.style.display='none'; }}
}}

function highlight(el, str) {{
    if (!str) {{
        el.childNodes.forEach(c => {{
            if (c.nodeType === 3) return;
            if (c.classList) c.classList.remove('hl');
            if (c.querySelectorAll) c.querySelectorAll('.hl').forEach(h => h.classList.remove('hl'));
        }});
        return;
    }}
    el.childNodes.forEach(c => {{
        if (c.nodeType === 3 && c.textContent.trim()) {{
            let idx = c.textContent.toLowerCase().indexOf(str);
            if (idx === -1) return;
            let span = document.createElement('span');
            let before = c.textContent.slice(0, idx);
            let match = c.textContent.slice(idx, idx + str.length);
            let after = c.textContent.slice(idx + str.length);
            let mk = document.createElement('mark');
            mk.className = 'hl';
            mk.textContent = match;
            if (before) span.appendChild(document.createTextNode(before));
            span.appendChild(mk);
            if (after) span.appendChild(document.createTextNode(after));
            el.replaceChild(span, c);
        }} else if (c.nodeType === 1) {{
            highlight(c, str);
        }}
    }});
}}

q.addEventListener('input', ()=>{{
    let v = q.value.toLowerCase();
    let rows = b.children;
    let vis = 0;
    document.querySelectorAll('.hl').forEach(h => {{
        let p = h.parentNode;
        let txt = document.createTextNode(h.textContent);
        p.replaceChild(txt, h);
        p.normalize();
    }});
    for (let r of rows) {{
        let ok = r.textContent.toLowerCase().includes(v);
        r.style.display = ok ? '' : 'none';
        if (ok && v) highlight(r, v);
        if (ok) vis++;
    }}
    cnt.textContent = vis + ' / ' + rows.length;
    vc.textContent = vis;
}});
</script>
</body></html>
"""

    return html


def _e(s):
    """Escape HTML entities."""
    s = str(s)
    s = s.replace("&", "&amp;")
    s = s.replace("<", "&lt;")
    s = s.replace(">", "&gt;")
    s = s.replace('"', "&quot;")
    return s


def main():
    args = sys.argv[1:]

    if len(args) >= 1:
        target = args[0]
    else:
        target = os.path.expanduser("~/mnt/nimo.loc/ComfyUI/models/loras")

    target = os.path.abspath(os.path.expanduser(target))

    if not os.path.isdir(target):
        print(f"Error: {target} is not a directory", file=sys.stderr)
        sys.exit(1)

    safetensors = sorted([
        f for f in os.listdir(target)
        if f.endswith(".safetensors") and not f.startswith("._")
    ])

    if not safetensors:
        print(f"No .safetensors files found in {target}", file=sys.stderr)
        sys.exit(1)

    loras = []
    for fname in safetensors:
        fpath = os.path.join(target, fname)
        meta = read_safetensors_meta(fpath)
        loras.append((fname, meta))

    # Look up source URLs with disk-backed cache
    cache = load_source_cache()
    sources = {}
    total = len(loras)
    for i, (fname, _) in enumerate(loras, 1):
        fpath = os.path.join(target, fname)
        url, label = lookup_source(fpath, fname, cache)
        if url:
            sources[fname] = (url, label)
        if i % 10 == 0 or i == total:
            print(f"  source lookup: {i}/{total}", file=sys.stderr)
    save_source_cache(cache)

    html = generate_html(loras, sources)

    if len(args) >= 2:
        out_path = os.path.abspath(os.path.expanduser(args[1]))
    else:
        out_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), "lora-info-sheet.html")

    with open(out_path, "w") as f:
        f.write(html)

    out_count = sum(1 for v in sources.values() if v[0])
    print(f"Generated {out_path} with {len(loras)} LoRA entries ({out_count} with source links)")


if __name__ == "__main__":
    main()
