import java.nio.charset.StandardCharsets;
import java.util.ArrayList;
import java.util.Arrays;
import java.util.Comparator;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import java.util.PriorityQueue;
import java.util.regex.Matcher;
import java.util.regex.Pattern;

/**
 * Byte-level BPE-tokenizer &mdash; mijlpaal M2.
 *
 * <p>Zet tekst om in tokennummers en terug. Alles wat daarvoor nodig is staat in het GGUF-bestand:
 * de woordenschat ({@code tokenizer.ggml.tokens}), de samenvoegregels ({@code tokenizer.ggml.merges})
 * en de tokensoorten ({@code tokenizer.ggml.token_type}).
 *
 * <p>Het gaat in vier stappen:
 * <ol>
 *   <li><b>Speciale tokens afsplitsen.</b> {@code <|im_start|>} en verwanten mogen nooit door BPE;
 *       ze worden letterlijk herkend en apart gezet.</li>
 *   <li><b>Voorsplitsen.</b> Een reguliere uitdrukking knipt de tekst in brokken, zodat een woord
 *       nooit over een spatie of leesteken heen samengevoegd wordt. Welke uitdrukking dat is, hangt
 *       af van {@code tokenizer.ggml.pre}.</li>
 *   <li><b>Bytes naar tekens.</b> Elke byte krijgt een eigen zichtbaar teken volgens de GPT-2-tabel.
 *       Daardoor kan élke bytereeks getokeniseerd worden, ook stukken die geen geldige tekst zijn.
 *       Een spatie wordt zo het teken {@code Ġ}.</li>
 *   <li><b>Samenvoegen.</b> Begin met losse bytes en voeg telkens het paar samen dat de laagste
 *       rang heeft in de samenvoeglijst, tot er niets meer samen kan.</li>
 * </ol>
 */
public final class Tokenizer {

    /** De tokensoorten zoals GGUF ze nummert. */
    public static final int NORMAL = 1, UNKNOWN = 2, CONTROL = 3, USER_DEFINED = 4, UNUSED = 5, BYTE = 6;

    // ---------------------------------------------------------------- byte-tabel van GPT-2

    private static final char[] B2C = new char[256];      // byte naar zichtbaar teken
    private static final int[]  C2B = new int[0x200];     // en terug
    static {
        Arrays.fill(C2B, -1);
        boolean[] direct = new boolean[256];
        for (int b = '!';  b <= '~';  b++) direct[b] = true;
        for (int b = 0xA1; b <= 0xAC; b++) direct[b] = true;
        for (int b = 0xAE; b <= 0xFF; b++) direct[b] = true;
        int extra = 0;
        for (int b = 0; b < 256; b++) {
            B2C[b] = direct[b] ? (char) b : (char) (256 + extra++);
            C2B[B2C[b]] = b;
        }
    }

    // ---------------------------------------------------------------- voorsplitsers

    /**
     * De reguliere uitdrukkingen per tokenizer-variant, overgenomen uit llama.cpp.
     * Het verschil tussen qwen2 en llama3 is subtiel maar wezenlijk: qwen2 heeft {@code \p{N}}
     * en knipt cijfers dus stuk voor stuk los, llama3 heeft {@code \p{N}{1,3}} en houdt ze
     * per drie bij elkaar. De qwen35-variant (Qwen3.5/3.8, 2026) voegt {@code \p{M}} toe:
     * combinerende tekens (accenten, Devanagari-klinkertekens) blijven bij hun letter in
     * plaats van als leesteken afgesplitst te worden.
     */
    private static final Map<String, String> PRE_TOKENIZERS = Map.of(
        "qwen2",
            "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])"
          + "|[^\\r\\n\\p{L}\\p{N}]?\\p{L}+|\\p{N}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*"
          + "|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+",
        "qwen35",
            "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])"
          + "|[^\\r\\n\\p{L}\\p{N}]?[\\p{L}\\p{M}]+|\\p{N}| ?[^\\s\\p{L}\\p{M}\\p{N}]+[\\r\\n]*"
          + "|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+",
        "llama3",
            "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])"
          + "|[^\\r\\n\\p{L}\\p{N}]?\\p{L}+|\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*"
          + "|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+",
        "default",
            "'s|'t|'re|'ve|'m|'ll|'d| ?\\p{L}+| ?\\p{N}+| ?[^\\s\\p{L}\\p{N}]+|\\s+(?!\\S)|\\s+"
    );

    // ---------------------------------------------------------------- toestand

    private record Merge(int rank, int result) {}

    private final String[] tokens;                 // byte-level tekst per id
    private final Map<String, Integer> idOf;
    private final int[] types;
    private final Map<Long, Merge> merges;         // (links << 32 | rechts) naar rang + resultaat
    private final java.util.BitSet mergeResults;   // ids die als samenvoegresultaat voorkomen
    private final Pattern preTokenizer;
    private final Pattern specials;                // null als er geen speciale tokens zijn
    private final String preName;

    public final int bosId, eosId;
    public final boolean addBos;

    // ---------------------------------------------------------------- opbouw

    public Tokenizer(Gguf g) {
        String model = g.has("tokenizer.ggml.model") ? g.getString("tokenizer.ggml.model") : "?";
        if (!model.equals("gpt2")) {
            throw new IllegalArgumentException("tokenizer.ggml.model = '" + model
                    + "'; deze klasse doet alleen byte-level BPE ('gpt2'). "
                    + "SentencePiece-modellen ('llama', 't5') werken anders.");
        }

        this.tokens = g.getStringArray("tokenizer.ggml.tokens");
        this.idOf = HashMap.newHashMap(tokens.length);
        for (int i = 0; i < tokens.length; i++) {
            idOf.putIfAbsent(tokens[i], i);         // bij dubbels wint de laagste id, zoals llama.cpp
        }

        this.types = new int[tokens.length];
        if (g.has("tokenizer.ggml.token_type")) {
            long[] t = g.getLongArray("tokenizer.ggml.token_type");
            for (int i = 0; i < types.length; i++) types[i] = i < t.length ? (int) t[i] : NORMAL;
        } else {
            Arrays.fill(types, NORMAL);
        }

        String[] mergeRules = g.getStringArray("tokenizer.ggml.merges");
        this.merges = HashMap.newHashMap(mergeRules.length);
        this.mergeResults = new java.util.BitSet(tokens.length);
        int skipped = 0;
        for (int rank = 0; rank < mergeRules.length; rank++) {
            // Een regel is "A B" met precies één scheidende spatie. De delen kunnen zelf
            // het teken Ġ bevatten, maar nooit een echte spatie -- die bestaat niet in
            // byte-level notatie.
            int sp = mergeRules[rank].indexOf(' ');
            if (sp < 0) { skipped++; continue; }
            Integer a = idOf.get(mergeRules[rank].substring(0, sp));
            Integer b = idOf.get(mergeRules[rank].substring(sp + 1));
            Integer ab = a == null || b == null ? null
                       : idOf.get(mergeRules[rank].substring(0, sp) + mergeRules[rank].substring(sp + 1));
            if (a == null || b == null || ab == null) { skipped++; continue; }
            merges.putIfAbsent(key(a, b), new Merge(rank, ab));
            mergeResults.set(ab);
        }
        this.skippedMerges = skipped;

        this.preName = g.has("tokenizer.ggml.pre") ? g.getString("tokenizer.ggml.pre") : "default";
        String rx = PRE_TOKENIZERS.get(preName);
        if (rx == null) {
            throw new IllegalArgumentException("onbekende voorsplitser '" + preName
                    + "'; bekend zijn " + PRE_TOKENIZERS.keySet()
                    + ". Raden is hier gevaarlijk: een verkeerde voorsplitser geeft geen foutmelding, "
                    + "alleen andere tokennummers.");
        }
        this.preTokenizer = Pattern.compile(rx, Pattern.UNICODE_CHARACTER_CLASS);

        // Speciale tokens: langste eerst, zodat een langer token wint van een korter dat erin zit.
        List<String> special = new ArrayList<>();
        for (int i = 0; i < tokens.length; i++) {
            if (types[i] == CONTROL || types[i] == USER_DEFINED) special.add(tokens[i]);
        }
        special.sort(Comparator.comparingInt(String::length).reversed());
        this.specials = special.isEmpty() ? null
                : Pattern.compile(String.join("|", special.stream().map(Pattern::quote).toList()));

        this.bosId = g.has("tokenizer.ggml.bos_token_id") ? g.getInt("tokenizer.ggml.bos_token_id") : -1;
        this.eosId = g.has("tokenizer.ggml.eos_token_id") ? g.getInt("tokenizer.ggml.eos_token_id") : -1;
        this.addBos = g.has("tokenizer.ggml.add_bos_token") && g.getBool("tokenizer.ggml.add_bos_token");
    }

    private final int skippedMerges;

    private static long key(int a, int b) {
        return (long) a << 32 | (b & 0xFFFFFFFFL);
    }

    // ---------------------------------------------------------------- coderen

    /** Codeert tekst, met herkenning van speciale tokens. */
    public int[] encode(String text) {
        return encode(text, true);
    }

    /**
     * De normalisatie die deze tokenizer-variant v&oacute;&oacute;r het voorsplitsen toepast.
     * De qwen35-tokenizer (Qwen3.5/3.8) schrijft NFC voor: een losse letter plus
     * combinerend accent wordt eerst het voorgevormde teken (e + ́ &rarr; &eacute;),
     * precies zoals de HuggingFace-referentie het doet. De qwen2-variant blijft
     * onaangeroerd: die is byte-voor-byte bewezen tegen llama.cpp, dat niet normaliseert.
     */
    public String normalize(String text) {
        return preName.equals("qwen35")
                ? java.text.Normalizer.normalize(text, java.text.Normalizer.Form.NFC)
                : text;
    }

    public int[] encode(String text, boolean parseSpecial) {
        IntList out = new IntList(text.length() / 3 + 8);
        if (parseSpecial && specials != null) {
            Matcher m = specials.matcher(text);
            int last = 0;
            while (m.find()) {
                if (m.start() > last) encodePlain(text.substring(last, m.start()), out);
                out.add(idOf.get(m.group()));
                last = m.end();
            }
            if (last < text.length()) encodePlain(text.substring(last), out);
        } else {
            encodePlain(text, out);
        }
        return out.toArray();
    }

    private void encodePlain(String text, IntList out) {
        text = normalize(text);
        Matcher m = preTokenizer.matcher(text);
        int covered = 0;
        while (m.find()) {
            if (m.start() > covered) {
                // De uitdrukkingen uit llama.cpp dekken alles af; gebeurt dit toch,
                // dan is er iets mis en willen we het weten in plaats van tekst te verliezen.
                throw new IllegalStateException("voorsplitser sloeg tekst over op positie " + covered);
            }
            if (!m.group().isEmpty()) bpe(m.group(), out);
            covered = m.end();
        }
        if (covered < text.length()) {
            throw new IllegalStateException("voorsplitser stopte op " + covered + " van " + text.length());
        }
    }

    /**
     * Het eigenlijke samenvoegen. Werkt met een dubbelgeschakelde lijst plus een prioriteitswachtrij,
     * zodat de kost O(n log n) blijft: bij een lange reeks spaties zou de naïeve aanpak
     * (telkens de hele reeks aflopen op zoek naar het beste paar) onwerkbaar traag worden.
     */
    private void bpe(String chunk, IntList out) {
        bpe(chunk, out, null);
    }

    /** Eén samenvoeging, zoals ze gebeurde. Alleen om te tonen wat er gebeurt. */
    public record MergeStep(int rank, String left, String right, String result) {}

    /**
     * Voert de samenvoeglus uit op één brok en geeft terug welke samenvoegingen er in welke
     * volgorde plaatsvonden. Bedoeld om uit te leggen en om fouten op te sporen; het rekenwerk
     * is exact hetzelfde als bij {@link #encode}.
     */
    public List<MergeStep> explain(String chunk) {
        List<MergeStep> trace = new ArrayList<>();
        bpe(chunk, new IntList(8), trace);
        return trace;
    }

    private void bpe(String chunk, IntList out, List<MergeStep> trace) {
        byte[] raw = chunk.getBytes(StandardCharsets.UTF_8);
        int n = raw.length;
        if (n == 0) return;

        int[] sym = new int[n];
        for (int i = 0; i < n; i++) {
            Integer id = idOf.get(String.valueOf(B2C[raw[i] & 0xFF]));
            if (id == null) {
                throw new IllegalStateException("byte 0x%02X heeft geen token in de woordenschat"
                        .formatted(raw[i] & 0xFF));
            }
            sym[i] = id;
        }
        if (n == 1) { out.add(sym[0]); return; }

        int[] prev = new int[n], next = new int[n];
        boolean[] alive = new boolean[n];
        for (int i = 0; i < n; i++) {
            prev[i] = i - 1;
            next[i] = i + 1 < n ? i + 1 : -1;
            alive[i] = true;
        }

        // rang, dan positie: bij gelijke rang wint de linkse, zoals de referentie doet
        PriorityQueue<int[]> pq = new PriorityQueue<>(
                Comparator.<int[]>comparingInt(c -> c[0]).thenComparingInt(c -> c[1]));
        for (int i = 0; i + 1 < n; i++) offer(pq, sym, i, i + 1);

        while (!pq.isEmpty()) {
            int[] c = pq.poll();               // {rang, links, rechts, resultaat, idLinks, idRechts}
            int l = c[1], r = c[2];
            if (!alive[l] || !alive[r] || next[l] != r) continue;   // paar is achterhaald
            if (sym[l] != c[4] || sym[r] != c[5]) continue;

            if (trace != null) {
                trace.add(new MergeStep(c[0], tokens[sym[l]], tokens[sym[r]], tokens[c[3]]));
            }
            sym[l] = c[3];
            alive[r] = false;
            next[l] = next[r];
            if (next[l] != -1) prev[next[l]] = l;

            if (prev[l] != -1) offer(pq, sym, prev[l], l);
            if (next[l] != -1) offer(pq, sym, l, next[l]);
        }

        for (int i = 0; i != -1; i = next[i]) out.add(sym[i]);
    }

    private void offer(PriorityQueue<int[]> pq, int[] sym, int l, int r) {
        Merge m = merges.get(key(sym[l], sym[r]));
        if (m != null) pq.add(new int[]{m.rank(), l, r, m.result(), sym[l], sym[r]});
    }

    // ---------------------------------------------------------------- decoderen

    /**
     * Zet tokennummers terug om in tekst. Let op: dit werkt alleen op een volledige reeks.
     * Eén token kan de helft van een teken bevatten &mdash; wie tijdens het genereren token per
     * token wil tonen, moet de bytes bufferen tot ze samen een geldig teken vormen.
     */
    public String decode(int[] ids) {
        return decode(ids, 0, ids.length);
    }

    public String decode(int[] ids, int from, int to) {
        byte[] buf = new byte[(to - from) * 4 + 8];
        int n = 0;
        for (int i = from; i < to; i++) {
            String t = tokens[ids[i]];
            if (n + t.length() > buf.length) buf = Arrays.copyOf(buf, Math.max(buf.length * 2, n + t.length()));
            for (int j = 0; j < t.length(); j++) {
                char ch = t.charAt(j);
                int b = ch < C2B.length ? C2B[ch] : -1;
                if (b < 0) {
                    throw new IllegalStateException("token %d bevat teken U+%04X dat niet in de byte-tabel staat"
                            .formatted(ids[i], (int) ch));
                }
                buf[n++] = (byte) b;
            }
        }
        return new String(buf, 0, n, StandardCharsets.UTF_8);
    }

    /** De ruwe bytes van één token &mdash; nodig om tijdens het genereren correct te kunnen streamen. */
    public byte[] tokenBytes(int id) {
        String t = tokens[id];
        byte[] b = new byte[t.length()];
        for (int j = 0; j < t.length(); j++) b[j] = (byte) C2B[t.charAt(j)];
        return b;
    }

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

    public int vocabSize()            { return tokens.length; }
    public String tokenText(int id)   { return tokens[id]; }
    public int typeOf(int id)         { return types[id]; }
    public boolean isSpecial(int id)  { return types[id] == CONTROL || types[id] == USER_DEFINED; }
    public int idOf(String text)      { return idOf.getOrDefault(text, -1); }
    public int mergeCount()           { return merges.size(); }
    public int skippedMerges()        { return skippedMerges; }

    /**
     * Waar als dit token het resultaat van minstens één samenvoegregel is. Een token van
     * meerdere tekens dat hier {@code false} geeft, kan de samenvoeglus per constructie
     * nooit bouwen — het staat wel in de woordenschat, maar er leidt geen pad naartoe
     * (qwen3.8 heeft er 201 van; qwen2.5 geen enkel).
     */
    public boolean isMergeResult(int id) { return mergeResults.get(id); }
    public String preTokenizerName()  { return preName; }

    /**
     * Hoeveel brokken de voorsplitser van deze tekst maakt, vóór er iets samengevoegd wordt.
     * Meer dan één brok betekent dat de tekst per definitie meerdere tokens oplevert.
     */
    public int chunkCount(String text) {
        Matcher m = preTokenizer.matcher(text);
        int n = 0;
        while (m.find()) if (!m.group().isEmpty()) n++;
        return n;
    }

    /** De brokken die de voorsplitser maakt &mdash; om te tonen waar een splitsing vandaan komt. */
    public List<String> chunks(String text) {
        Matcher m = preTokenizer.matcher(text);
        List<String> out = new ArrayList<>();
        while (m.find()) if (!m.group().isEmpty()) out.add(m.group());
        return out;
    }

    /** De tokens als leesbare tekst, met Ġ voor spatie &mdash; handig om een splitsing te tonen. */
    public String[] pieces(int[] ids) {
        String[] p = new String[ids.length];
        for (int i = 0; i < ids.length; i++) p[i] = tokens[ids[i]];
        return p;
    }

    /** Eenvoudige groeiende int-lijst; java.util.List<Integer> zou hier onnodig veel doosjes maken. */
    static final class IntList {
        int[] a;
        int n;
        IntList(int cap) { a = new int[Math.max(cap, 8)]; }
        void add(int v) {
            if (n == a.length) a = Arrays.copyOf(a, a.length * 2);
            a[n++] = v;
        }
        int[] toArray() { return Arrays.copyOf(a, n); }
    }
}
