import java.io.IOException;
import java.lang.foreign.Arena;
import java.lang.foreign.MemorySegment;
import java.lang.foreign.ValueLayout;
import java.nio.ByteOrder;
import java.nio.channels.FileChannel;
import java.nio.charset.StandardCharsets;
import java.nio.file.Files;
import java.nio.file.Path;
import java.nio.file.StandardOpenOption;
import java.util.ArrayList;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import java.util.Set;

/**
 * Lezer voor GGUF-bestanden &mdash; mijlpaal M1.
 *
 * <p>Een GGUF-bestand bestaat uit vier delen die netjes na elkaar staan:
 * <ol>
 *   <li>een header van 24 bytes: magisch getal, versie, aantal tensoren, aantal metadata-sleutels;</li>
 *   <li>de metadata: sleutel-waardeparen met alle hyperparameters én de volledige tokenizer;</li>
 *   <li>het tensorregister: naam, vorm, type en positie van elke tensor;</li>
 *   <li>opvulling tot de uitlijningsgrens, en dan de ruwe gewichten.</li>
 * </ol>
 *
 * <p>Alles is klein-endisch. Het bestand wordt in zijn geheel gemapt met
 * {@link FileChannel#map(FileChannel.MapMode, long, long, Arena)}, wat sinds JDK 22 een
 * {@link MemorySegment} teruggeeft in plaats van een {@code MappedByteBuffer}. Dat is geen detail:
 * een MappedByteBuffer loopt vast op 2 GB, en modellen zijn groter.
 *
 * <p>Deze klasse leest <em>alleen de structuur</em>. De gewichten blijven op schijf tot iemand ze
 * werkelijk aanraakt via {@link Tensor#data()}.
 */
public final class Gguf implements AutoCloseable {

    /** De vier bytes 'G','G','U','F', klein-endisch gelezen als int. */
    public static final int MAGIC = 0x46554747;

    /** Uitlijning van het datablok wanneer het bestand niets anders opgeeft. */
    public static final long DEFAULT_ALIGNMENT = 32;

    // Alle layouts expliciet klein-endisch: op ARM en x86 is dat toevallig de eigen volgorde,
    // maar dat mag je niet als gegeven beschouwen.
    private static final ValueLayout.OfShort  I16 = ValueLayout.JAVA_SHORT_UNALIGNED.withOrder(ByteOrder.LITTLE_ENDIAN);
    private static final ValueLayout.OfInt    I32 = ValueLayout.JAVA_INT_UNALIGNED.withOrder(ByteOrder.LITTLE_ENDIAN);
    private static final ValueLayout.OfLong   I64 = ValueLayout.JAVA_LONG_UNALIGNED.withOrder(ByteOrder.LITTLE_ENDIAN);
    private static final ValueLayout.OfFloat  F32 = ValueLayout.JAVA_FLOAT_UNALIGNED.withOrder(ByteOrder.LITTLE_ENDIAN);
    private static final ValueLayout.OfDouble F64 = ValueLayout.JAVA_DOUBLE_UNALIGNED.withOrder(ByteOrder.LITTLE_ENDIAN);

    /** De dertien waardetypes die in de metadata kunnen voorkomen. */
    public enum Kind {
        UINT8(0), INT8(1), UINT16(2), INT16(3), UINT32(4), INT32(5), FLOAT32(6),
        BOOL(7), STRING(8), ARRAY(9), UINT64(10), INT64(11), FLOAT64(12);

        public final int id;
        Kind(int id) { this.id = id; }

        static Kind byId(int id) {
            for (Kind k : values()) if (k.id == id) return k;
            throw new IllegalArgumentException("onbekend metadata-type " + id);
        }
    }

    /**
     * Eén metadata-waarde, met het type zoals het in het bestand stond.
     *
     * @param kind het gedeclareerde type
     * @param elem bij een array: het type van de elementen, anders {@code null}
     * @param data {@code Long}/{@code Double}/{@code Boolean}/{@code String}, of bij een array
     *             een {@code String[]}, {@code long[]}, {@code double[]} of {@code boolean[]}
     */
    public record Value(Kind kind, Kind elem, Object data) {

        public boolean isArray() { return kind == Kind.ARRAY; }

        public int length() {
            return switch (data) {
                case String[] a  -> a.length;
                case long[] a    -> a.length;
                case double[] a  -> a.length;
                case boolean[] a -> a.length;
                default          -> 1;
            };
        }
    }

    private final Path path;
    private final Arena arena;
    private final MemorySegment file;
    private final long fileSize;
    private final int version;
    private final long alignment;
    private final long dataOffset;
    private final Map<String, Value> metadata;
    private final Map<String, Tensor> tensors;

    private long p;   // leescursor, alleen in gebruik tijdens het ontleden

    // ---------------------------------------------------------------- openen

    public static Gguf open(Path path) throws IOException {
        return new Gguf(path);
    }

    private Gguf(Path path) throws IOException {
        this.path = path;
        this.fileSize = Files.size(path);
        if (fileSize < 24) {
            throw new IOException("bestand is te klein om een GGUF-header te bevatten: " + fileSize + " bytes");
        }

        this.arena = Arena.ofShared();
        try (FileChannel ch = FileChannel.open(path, StandardOpenOption.READ)) {
            // De mapping hoort bij de arena, niet bij het kanaal: ze blijft geldig na deze try.
            this.file = ch.map(FileChannel.MapMode.READ_ONLY, 0, fileSize, arena);
        } catch (IOException | RuntimeException e) {
            arena.close();
            throw e;
        }

        try {
            int magic = u32();
            if (magic != MAGIC) {
                throw new IOException("geen GGUF-bestand (magisch getal 0x%08X, verwacht 0x%08X)"
                        .formatted(magic, MAGIC));
            }
            this.version = u32();
            if (version < 2 || version > 3) {
                // v1 gebruikte 32-bits lengtes voor strings en arrays; die ontleding is anders.
                throw new IOException("GGUF-versie " + version + " wordt niet ondersteund (verwacht 2 of 3)");
            }
            long tensorCount = u64();
            long kvCount = u64();
            if (tensorCount < 0 || kvCount < 0 || tensorCount > 1_000_000 || kvCount > 1_000_000) {
                throw new IOException("onwaarschijnlijke aantallen in de header: "
                        + tensorCount + " tensoren, " + kvCount + " metadata-sleutels");
            }

            this.metadata = new LinkedHashMap<>();
            for (long i = 0; i < kvCount; i++) {
                String key = str();
                Value val = value();
                if (metadata.put(key, val) != null) {
                    throw new IOException("dubbele metadata-sleutel: " + key);
                }
            }

            this.alignment = metadata.containsKey("general.alignment")
                    ? getLong("general.alignment")
                    : DEFAULT_ALIGNMENT;
            if (alignment <= 0 || Long.bitCount(alignment) != 1) {
                throw new IOException("general.alignment moet een macht van twee zijn, niet " + alignment);
            }

            // Eerst alle registeringangen lezen; de absolute posities kennen we pas
            // wanneer we weten waar het datablok begint.
            record Raw(String name, long[] dims, GgmlType type, long offset) {}
            List<Raw> raws = new ArrayList<>((int) tensorCount);
            for (long i = 0; i < tensorCount; i++) {
                String name = str();
                int nDims = u32();
                if (nDims < 1 || nDims > 4) {
                    throw new IOException("tensor '" + name + "' heeft " + nDims + " dimensies; ggml staat 1 tot 4 toe");
                }
                long[] dims = new long[nDims];
                for (int d = 0; d < nDims; d++) {
                    dims[d] = u64();
                    if (dims[d] <= 0) throw new IOException("tensor '" + name + "' heeft dimensie " + dims[d]);
                }
                int typeId = u32();
                if (!GgmlType.isKnown(typeId)) {
                    throw new IOException("tensor '" + name + "': " + new IllegalArgumentException(
                            "ggml-type " + typeId + " is onbekend of nog niet ondersteund").getMessage());
                }
                GgmlType type = GgmlType.byId(typeId);
                long offset = u64();
                raws.add(new Raw(name, dims, type, offset));
            }

            // Het datablok begint op de eerstvolgende uitlijningsgrens na het register.
            this.dataOffset = align(p, alignment);

            this.tensors = new LinkedHashMap<>();
            for (Raw r : raws) {
                long elements = 1;
                for (long d : r.dims()) elements *= d;

                if (r.dims()[0] % r.type().blockSize != 0) {
                    throw new IOException("tensor '" + r.name() + "': rijlengte " + r.dims()[0]
                            + " is niet deelbaar door de blokgrootte " + r.type().blockSize
                            + " van " + r.type());
                }
                long bytes = r.type().byteSize(elements);
                long abs = dataOffset + r.offset();
                if (abs < dataOffset || abs + bytes > fileSize) {
                    throw new IOException("tensor '" + r.name() + "' valt buiten het bestand: "
                            + abs + " + " + bytes + " > " + fileSize);
                }
                Tensor t = new Tensor(r.name(), r.dims(), r.type(), r.offset(), bytes,
                        file.asSlice(abs, bytes));
                if (tensors.put(r.name(), t) != null) {
                    throw new IOException("dubbele tensornaam: " + r.name());
                }
            }
        } catch (IOException | RuntimeException e) {
            arena.close();
            throw e;
        }
    }

    @Override
    public void close() {
        arena.close();
    }

    private static long align(long value, long alignment) {
        return (value + alignment - 1) / alignment * alignment;
    }

    // ---------------------------------------------------------------- ruwe lezers

    private byte u8()    { byte v = file.get(ValueLayout.JAVA_BYTE, p); p += 1; return v; }
    private short u16()  { short v = file.get(I16, p); p += 2; return v; }
    private int u32()    { int v = file.get(I32, p); p += 4; return v; }
    private long u64()   { long v = file.get(I64, p); p += 8; return v; }
    private float f32()  { float v = file.get(F32, p); p += 4; return v; }
    private double f64() { double v = file.get(F64, p); p += 8; return v; }

    /** Een GGUF-string: 64-bits lengte, dan zoveel UTF-8-bytes. Niet nul-afgesloten. */
    private String str() throws IOException {
        long n = u64();
        if (n < 0 || n > 1 << 26) {
            throw new IOException("onwaarschijnlijke stringlengte " + n + " op positie " + (p - 8));
        }
        if (p + n > fileSize) {
            throw new IOException("string van %d bytes op positie %d loopt voorbij het einde van het bestand (%d); afgekapt bestand?"
                    .formatted(n, p, fileSize));
        }
        byte[] b = new byte[(int) n];
        MemorySegment.copy(file, ValueLayout.JAVA_BYTE, p, b, 0, (int) n);
        p += n;
        return new String(b, StandardCharsets.UTF_8);
    }

    private Value value() throws IOException {
        Kind kind = Kind.byId(u32());
        if (kind != Kind.ARRAY) {
            return new Value(kind, null, scalar(kind));
        }
        Kind elem = Kind.byId(u32());
        long n = u64();
        if (n < 0 || n > 1 << 26) {
            throw new IOException("onwaarschijnlijke arraylengte " + n);
        }
        int len = (int) n;
        // Per elementtype een echte primitieve array: bij 128 000 tokens scheelt dat
        // tientallen megabytes aan doosjes.
        Object data = switch (elem) {
            case STRING -> {
                String[] a = new String[len];
                for (int i = 0; i < len; i++) a[i] = str();
                yield a;
            }
            case FLOAT32, FLOAT64 -> {
                double[] a = new double[len];
                for (int i = 0; i < len; i++) a[i] = (elem == Kind.FLOAT32) ? f32() : f64();
                yield a;
            }
            case BOOL -> {
                boolean[] a = new boolean[len];
                for (int i = 0; i < len; i++) a[i] = u8() != 0;
                yield a;
            }
            case ARRAY -> throw new IOException("geneste arrays komen in GGUF niet voor");
            default -> {
                long[] a = new long[len];
                for (int i = 0; i < len; i++) a[i] = integral(elem);
                yield a;
            }
        };
        return new Value(Kind.ARRAY, elem, data);
    }

    private Object scalar(Kind kind) throws IOException {
        return switch (kind) {
            case STRING  -> str();
            case FLOAT32 -> (double) f32();
            case FLOAT64 -> f64();
            case BOOL    -> u8() != 0;
            case ARRAY   -> throw new IllegalStateException("onbereikbaar");
            default      -> integral(kind);
        };
    }

    /**
     * Leest een geheel getal en verbreedt het naar long. De niet-ondertekende types worden
     * netjes uitgebreid, zodat een UINT32 van 3 miljard niet als negatief getal terugkomt.
     */
    private long integral(Kind kind) {
        return switch (kind) {
            case UINT8  -> Byte.toUnsignedInt(u8());
            case INT8   -> u8();
            case UINT16 -> Short.toUnsignedInt(u16());
            case INT16  -> u16();
            case UINT32 -> Integer.toUnsignedLong(u32());
            case INT32  -> u32();
            case UINT64, INT64 -> u64();   // een echte uint64 boven 2^63 komt in de praktijk niet voor
            default -> throw new IllegalArgumentException(kind + " is geen geheel getal");
        };
    }

    // ---------------------------------------------------------------- toegang

    public Path path()                    { return path; }
    public long fileSize()                { return fileSize; }
    public int version()                  { return version; }
    public long alignment()               { return alignment; }
    public long dataOffset()              { return dataOffset; }
    public Map<String, Value> metadata()  { return java.util.Collections.unmodifiableMap(metadata); }
    public Map<String, Tensor> tensors()  { return java.util.Collections.unmodifiableMap(tensors); }
    public MemorySegment segment()        { return file; }

    public boolean has(String key) {
        return metadata.containsKey(key);
    }

    public Value value(String key) {
        Value v = metadata.get(key);
        if (v == null) throw new NoSuchElementException(key);
        return v;
    }

    public Tensor tensor(String name) {
        Tensor t = tensors.get(name);
        if (t == null) throw new NoSuchElementException("tensor " + name);
        return t;
    }

    public Tensor tensorOrNull(String name) {
        return tensors.get(name);
    }

    public String getString(String key) {
        Object d = value(key).data();
        if (!(d instanceof String s)) throw new IllegalArgumentException(key + " is geen string");
        return s;
    }

    public long getLong(String key) {
        Object d = value(key).data();
        if (d instanceof Long l) return l;
        if (d instanceof Double x) return (long) (double) x;
        throw new IllegalArgumentException(key + " is geen getal");
    }

    public int getInt(String key) {
        long v = getLong(key);
        if (v < Integer.MIN_VALUE || v > Integer.MAX_VALUE) {
            throw new ArithmeticException(key + " past niet in een int: " + v);
        }
        return (int) v;
    }

    public int getInt(String key, int fallback) {
        return has(key) ? getInt(key) : fallback;
    }

    public float getFloat(String key) {
        Object d = value(key).data();
        if (d instanceof Double x) return (float) (double) x;
        if (d instanceof Long l) return l;
        throw new IllegalArgumentException(key + " is geen kommagetal");
    }

    public float getFloat(String key, float fallback) {
        return has(key) ? getFloat(key) : fallback;
    }

    public boolean getBool(String key) {
        Object d = value(key).data();
        if (!(d instanceof Boolean b)) throw new IllegalArgumentException(key + " is geen booleaanse waarde");
        return b;
    }

    public String[] getStringArray(String key) {
        Object d = value(key).data();
        if (!(d instanceof String[] a)) throw new IllegalArgumentException(key + " is geen stringarray");
        return a;
    }

    public long[] getLongArray(String key) {
        Object d = value(key).data();
        if (!(d instanceof long[] a)) throw new IllegalArgumentException(key + " is geen array van gehele getallen");
        return a;
    }

    public double[] getDoubleArray(String key) {
        Object d = value(key).data();
        if (!(d instanceof double[] a)) throw new IllegalArgumentException(key + " is geen array van kommagetallen");
        return a;
    }

    /** De architectuur, bv. {@code llama} of {@code qwen2}. Bepaalt het voorvoegsel van de hyperparameters. */
    public String architecture() {
        return has("general.architecture") ? getString("general.architecture") : "onbekend";
    }

    /** Bouwt een architectuurgebonden sleutel: {@code arch("block_count")} wordt {@code llama.block_count}. */
    public String arch(String suffix) {
        return architecture() + "." + suffix;
    }

    public long totalTensorBytes() {
        long n = 0;
        for (Tensor t : tensors.values()) n += t.byteSize();
        return n;
    }

    public long totalParameters() {
        long n = 0;
        for (Tensor t : tensors.values()) n += t.elements();
        return n;
    }

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

    /**
     * Controleert het bestand op inwendige tegenstrijdigheden en geeft een lijst van klachten terug.
     * Een lege lijst betekent dat register, uitlijning en bestandsgrootte samen kloppen.
     */
    public List<String> verify() {
        List<String> problems = new ArrayList<>();

        if (dataOffset % alignment != 0) {
            problems.add("het datablok begint op %d, niet uitgelijnd op %d".formatted(dataOffset, alignment));
        }

        List<Tensor> byOffset = new ArrayList<>(tensors.values());
        byOffset.sort((a, b) -> Long.compare(a.offsetInData(), b.offsetInData()));

        long expectedNext = 0;
        Tensor previous = null;
        for (Tensor t : byOffset) {
            if (t.offsetInData() % alignment != 0) {
                problems.add("%s begint op %d, niet uitgelijnd op %d"
                        .formatted(t.name(), t.offsetInData(), alignment));
            }
            if (t.offsetInData() < expectedNext) {
                problems.add("%s overlapt met %s".formatted(t.name(),
                        previous == null ? "het begin van het datablok" : previous.name()));
            }
            expectedNext = t.offsetInData() + t.byteSize();
            previous = t;
        }

        long endOfData = dataOffset + expectedNext;
        if (endOfData > fileSize) {
            problems.add("de tensoren lopen tot %d, voorbij het einde van het bestand (%d)"
                    .formatted(endOfData, fileSize));
        } else {
            long slack = fileSize - endOfData;
            if (slack >= alignment) {
                problems.add("%d bytes ongebruikt aan het einde van het bestand".formatted(slack));
            }
        }

        return problems;
    }

    /**
     * Controleert of de tensoren aanwezig zijn die een decoder van deze architectuur nodig heeft.
     * Ontbrekende namen zijn niet altijd een fout: modellen met gebonden inbeddingen hebben geen
     * {@code output.weight} en hergebruiken {@code token_embd.weight}.
     */
    public List<String> missingTensors() {
        List<String> missing = new ArrayList<>();
        if (!tensors.containsKey("token_embd.weight")) missing.add("token_embd.weight");
        if (!tensors.containsKey("output_norm.weight")) missing.add("output_norm.weight");

        if (!has(arch("block_count"))) return missing;
        int layers = getInt(arch("block_count"));
        Set<String> perLayer = Set.of(
                "attn_norm.weight", "attn_q.weight", "attn_k.weight", "attn_v.weight",
                "attn_output.weight", "ffn_norm.weight", "ffn_gate.weight",
                "ffn_up.weight", "ffn_down.weight");
        for (int l = 0; l < layers; l++) {
            for (String suffix : perLayer) {
                String name = "blk." + l + "." + suffix;
                if (!tensors.containsKey(name)) missing.add(name);
            }
        }
        return missing;
    }

    /** Waar of de eindprojectie de inbeddingsmatrix hergebruikt. */
    public boolean hasTiedEmbeddings() {
        return !tensors.containsKey("output.weight") && tensors.containsKey("token_embd.weight");
    }

    public static final class NoSuchElementException extends java.util.NoSuchElementException {
        NoSuchElementException(String key) {
            super("geen metadata of tensor met sleutel '" + key + "'");
        }
    }
}
