#!/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


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]


def generate_html(loras):
    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>"

        rows.append((fname_no_ext, tool, base, triggers))

    num_files = len(rows)

    table_rows = ""
    for fname, tool, base, triggers in rows:
        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"    </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}}
.c0{{width:26%}}.c1{{width:9%}}.c2{{width:13%}}.c3{{width:52%}}
@media(max-width:768px){{.c1,.c2,td:nth-child(2),td:nth-child(3){{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>
</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))

    html = generate_html(loras)

    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)

    print(f"Generated {out_path} with {len(loras)} LoRA entries")


if __name__ == "__main__":
    main()
