import java.nio.file.Files;
import java.nio.file.Path;
import java.util.ArrayList;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;

/**
 * Drukt de structuur van een GGUF-bestand af, zoals {@code gguf_dump.py} uit llama.cpp.
 *
 * <p>Dit is het toetsingsgereedschap voor mijlpaal M1: als de lagen, koppen en woordenschat
 * hier overeenkomen met de modelkaart of met de uitvoer van llama.cpp, staat het fundament recht.
 *
 * <pre>
 *   java src/GgufDump.java model.gguf
 *   java src/GgufDump.java model.gguf --brief      alleen samenvatting en controles
 *   java src/GgufDump.java model.gguf --tokens 20  toon de eerste 20 tokens
 * </pre>
 */
public final class GgufDump {

    public static void main(String[] args) throws Exception {
        if (args.length == 0) {
            System.err.println("""
                    gebruik: java src/GgufDump.java <model.gguf> [opties]

                      --brief        laat het tensorregister weg
                      --tokens N     toon de eerste N tokens van de tokenizer (standaard 0)
                      --meta-only    alleen de metadata
                    """);
            System.exit(2);
        }

        Path path = Path.of(args[0]);
        if (!Files.isReadable(path)) {
            System.err.println("kan niet lezen: " + path.toAbsolutePath());
            System.exit(1);
        }

        boolean brief = false, metaOnly = false;
        int showTokens = 0;
        for (int i = 1; i < args.length; i++) {
            switch (args[i]) {
                case "--brief"     -> brief = true;
                case "--meta-only" -> metaOnly = true;
                case "--tokens"    -> showTokens = Integer.parseInt(args[++i]);
                default -> {
                    System.err.println("onbekende optie: " + args[i]);
                    System.exit(2);
                }
            }
        }

        long t0 = System.nanoTime();
        try (Gguf g = Gguf.open(path)) {
            long ms = (System.nanoTime() - t0) / 1_000_000;

            header("BESTAND");
            kv("pad", path.toAbsolutePath().toString());
            kv("grootte", bytes(g.fileSize()));
            kv("GGUF-versie", String.valueOf(g.version()));
            kv("uitlijning", g.alignment() + " bytes");
            kv("datablok begint op", g.dataOffset() + " (0x%X)".formatted(g.dataOffset()));
            kv("ontleed in", ms + " ms");

            printModelSummary(g);
            printMetadata(g, showTokens);
            if (!brief && !metaOnly) printTensors(g);
            if (!metaOnly) printChecks(g);
        }
    }

    // ---------------------------------------------------------------- samenvatting

    private static void printModelSummary(Gguf g) {
        header("MODEL");
        String a = g.architecture();
        kv("architectuur", a);
        if (g.has("general.name")) kv("naam", g.getString("general.name"));

        Map<String, String> rows = new LinkedHashMap<>();
        put(rows, g, "lagen",                 a + ".block_count");
        put(rows, g, "inbeddingsdimensie",    a + ".embedding_length");
        put(rows, g, "aandachtskoppen",       a + ".attention.head_count");
        put(rows, g, "kv-koppen",             a + ".attention.head_count_kv");
        put(rows, g, "feed-forward-dimensie", a + ".feed_forward_length");
        put(rows, g, "contextlengte",         a + ".context_length");
        put(rows, g, "RoPE-basis",            a + ".rope.freq_base");
        put(rows, g, "RMS-epsilon",           a + ".attention.layer_norm_rms_epsilon");
        rows.forEach(GgufDump::kv);

        // Afgeleide waarden: die staan niet in het bestand maar zeggen het meest.
        String heads = a + ".attention.head_count";
        String embd = a + ".embedding_length";
        if (g.has(heads) && g.has(embd)) {
            int nHeads = g.getInt(heads);
            int dim = g.getInt(embd);
            if (nHeads > 0 && dim % nHeads == 0) {
                kv("kopdimensie", String.valueOf(dim / nHeads) + "  (afgeleid)");
            }
            if (g.has(a + ".attention.head_count_kv")
                    && !(g.value(a + ".attention.head_count_kv").data() instanceof long[])) {
                int kv = g.getInt(a + ".attention.head_count_kv");
                if (kv > 0 && nHeads % kv == 0) {
                    int mul = nHeads / kv;
                    kv("GQA-groepering", mul == 1
                            ? "geen, elke kop heeft eigen K en V  (afgeleid)"
                            : mul + " vraagkoppen per kv-kop  (afgeleid)");
                }
            }
        }
        if (g.has("tokenizer.ggml.tokens")) {
            kv("woordenschat", String.valueOf(g.getStringArray("tokenizer.ggml.tokens").length));
        }
        if (g.has("tokenizer.ggml.model")) {
            kv("tokenizer", g.getString("tokenizer.ggml.model"));
        }

        if (!g.tensors().isEmpty()) {
            kv("tensoren", String.valueOf(g.tensors().size()));
            long params = g.totalParameters();
            String scaled = params >= 1_000_000_000L ? "  (%.2f B)".formatted(params / 1e9)
                          : params >= 1_000_000L     ? "  (%.0f M)".formatted(params / 1e6)
                          : "";
            kv("parameters", "%,d%s".formatted(params, scaled));
            kv("gewichten samen", bytes(g.totalTensorBytes()));
            kv("gemiddeld per gewicht", "%.2f bit".formatted(g.totalTensorBytes() * 8.0 / params));
            kv("gebonden inbeddingen", g.hasTiedEmbeddings() ? "ja, output.weight ontbreekt" : "nee");
        }
    }

    private static void put(Map<String, String> rows, Gguf g, String label, String key) {
        if (!g.has(key)) return;
        Gguf.Value v = g.value(key);
        // hybride architecturen (qwen35) geven kv-koppen als array per laag: 0 = recurrent
        if (v.data() instanceof long[] arr) {
            StringBuilder sb = new StringBuilder("per laag: [");
            for (int i = 0; i < arr.length && i < 26; i++) sb.append(i > 0 ? ", " : "").append(arr[i]);
            rows.put(label, sb.append(arr.length > 26 ? ", …]" : "]").toString());
            return;
        }
        rows.put(label, v.data() instanceof Double d ? trimFloat(d) : String.valueOf(v.data()));
    }

    // ---------------------------------------------------------------- metadata

    private static void printMetadata(Gguf g, int showTokens) {
        header("METADATA  (" + g.metadata().size() + " sleutels)");
        int w = g.metadata().keySet().stream().mapToInt(String::length).max().orElse(10);
        w = Math.min(w, 44);
        for (Map.Entry<String, Gguf.Value> e : g.metadata().entrySet()) {
            Gguf.Value v = e.getValue();
            String type = v.isArray() ? v.elem() + "[" + v.length() + "]" : v.kind().toString();
            System.out.printf("  %-" + w + "s  %-14s  %s%n", e.getKey(), type, preview(v));
        }

        if (showTokens > 0 && g.has("tokenizer.ggml.tokens")) {
            String[] toks = g.getStringArray("tokenizer.ggml.tokens");
            long[] types = g.has("tokenizer.ggml.token_type") ? g.getLongArray("tokenizer.ggml.token_type") : null;
            header("TOKENS  (eerste " + Math.min(showTokens, toks.length) + " van " + toks.length + ")");
            for (int i = 0; i < Math.min(showTokens, toks.length); i++) {
                System.out.printf("  %6d  %-24s %s%n", i, quote(toks[i]),
                        types != null && i < types.length ? "type=" + types[i] : "");
            }
        }
    }

    private static String preview(Gguf.Value v) {
        Object d = v.data();
        return switch (d) {
            case String s    -> quote(s);
            case Double x    -> trimFloat(x);
            case Boolean b   -> b.toString();
            case String[] a  -> joinArray(a.length, i -> quote(a[i]));
            case long[] a    -> joinArray(a.length, i -> String.valueOf(a[i]));
            case double[] a  -> joinArray(a.length, i -> trimFloat(a[i]));
            case boolean[] a -> joinArray(a.length, i -> String.valueOf(a[i]));
            default          -> String.valueOf(d);
        };
    }

    private interface Item { String at(int i); }

    private static String joinArray(int n, Item f) {
        int show = Math.min(n, 6);
        StringBuilder sb = new StringBuilder("[");
        for (int i = 0; i < show; i++) {
            if (i > 0) sb.append(", ");
            sb.append(f.at(i));
        }
        if (n > show) sb.append(", ... ").append(n - show).append(" meer");
        return sb.append(']').toString();
    }

    private static String quote(String s) {
        String t = s.replace("\n", "\\n").replace("\t", "\\t").replace("\r", "\\r");
        if (t.length() > 48) t = t.substring(0, 45) + "...";
        return '"' + t + '"';
    }

    private static String trimFloat(double d) {
        if (d == Math.rint(d) && Math.abs(d) < 1e15) return String.valueOf((long) d);
        return String.valueOf((float) d);
    }

    // ---------------------------------------------------------------- tensoren

    private static void printTensors(Gguf g) {
        header("TENSOREN  (" + g.tensors().size() + ")");
        int w = g.tensors().keySet().stream().mapToInt(String::length).max().orElse(20);
        System.out.printf("  %-" + w + "s  %-6s %-18s %14s %12s%n", "naam", "type", "vorm", "positie", "bytes");
        for (Tensor t : g.tensors().values()) {
            System.out.printf("  %-" + w + "s  %-6s %-18s %14d %12d%n",
                    t.name(), t.type(), t.shape(), t.offsetInData(), t.byteSize());
        }

        // Verdeling per type: laat meteen zien welke kwantisatie het model werkelijk gebruikt.
        Map<GgmlType, long[]> perType = new LinkedHashMap<>();   // [aantal, gewichten, bytes]
        for (Tensor t : g.tensors().values()) {
            long[] acc = perType.computeIfAbsent(t.type(), k -> new long[3]);
            acc[0]++;
            acc[1] += t.elements();
            acc[2] += t.byteSize();
        }
        header("VERDELING PER TYPE");
        System.out.printf("  %-6s %8s %16s %14s %10s%n", "type", "tensors", "gewichten", "bytes", "bit/gew.");
        perType.forEach((type, acc) -> System.out.printf("  %-6s %8d %16d %14s %10.2f%n",
                type, acc[0], acc[1], bytes(acc[2]), acc[2] * 8.0 / acc[1]));
    }

    // ---------------------------------------------------------------- controles

    private static void printChecks(Gguf g) {
        header("CONTROLES");

        List<String> problems = g.verify();
        if (problems.isEmpty()) {
            ok("register, uitlijning en bestandsgrootte kloppen onderling");
        } else {
            for (String p : problems) fail(p);
        }

        List<String> missing = g.missingTensors();
        if (missing.isEmpty()) {
            ok("alle tensoren aanwezig die een decoder van deze architectuur nodig heeft");
        } else if (missing.size() <= 6) {
            for (String m : missing) fail("ontbrekende tensor: " + m);
        } else {
            fail(missing.size() + " ontbrekende tensoren, o.a. " + String.join(", ", missing.subList(0, 5)));
        }

        // Kruiscontrole: komt de vorm van de tensoren overeen met wat de metadata beweert?
        String a = g.architecture();
        if (g.has(a + ".embedding_length")) {
            int dim = g.getInt(a + ".embedding_length");
            Tensor emb = g.tensorOrNull("token_embd.weight");
            if (emb != null) {
                if (emb.dims()[0] == dim) {
                    ok("token_embd.weight klopt met " + a + ".embedding_length (" + dim + ")");
                } else {
                    fail("token_embd.weight heeft rijlengte " + emb.dims()[0]
                            + " maar de metadata zegt " + dim);
                }
                if (g.has("tokenizer.ggml.tokens")) {
                    int vocab = g.getStringArray("tokenizer.ggml.tokens").length;
                    if (emb.dims().length > 1 && emb.dims()[1] == vocab) {
                        ok("token_embd.weight klopt met de woordenschat (" + vocab + ")");
                    } else if (emb.dims().length > 1) {
                        fail("token_embd.weight heeft " + emb.dims()[1]
                                + " rijen maar de tokenizer telt " + vocab + " tokens");
                    }
                }
            }
        }
        if (g.has(a + ".block_count")) {
            int layers = g.getInt(a + ".block_count");
            long found = g.tensors().keySet().stream()
                    .filter(n -> n.startsWith("blk."))
                    .map(n -> n.split("\\.")[1])
                    .distinct().count();
            if (found == layers) {
                ok("het bestand bevat precies " + layers + " lagen");
            } else {
                fail("de metadata zegt " + layers + " lagen, maar er zijn tensoren voor " + found);
            }
        }
    }

    // ---------------------------------------------------------------- opmaak

    private static void header(String title) {
        System.out.println();
        System.out.println("== " + title + " " + "=".repeat(Math.max(0, 62 - title.length())));
    }

    private static void kv(String k, String v) {
        System.out.printf("  %-24s %s%n", k, v);
    }

    private static void ok(String msg)   { System.out.println("  [ok]   " + msg); }
    private static void fail(String msg) { System.out.println("  [FOUT] " + msg); }

    static String bytes(long n) {
        if (n < 1024) return n + " B";
        String[] u = {"KiB", "MiB", "GiB", "TiB"};
        double d = n;
        int i = -1;
        while (d >= 1024 && i < u.length - 1) { d /= 1024; i++; }
        return "%.2f %s".formatted(d, u[i]);
    }
}
