import java.io.BufferedReader;
import java.io.InputStreamReader;
import java.nio.charset.StandardCharsets;
import java.nio.file.Files;
import java.nio.file.Path;

/**
 * De interactieve chat &mdash; mijlpaal M6/M7, het sluitstuk van het project.
 *
 * <pre>
 *   java --add-modules jdk.incubator.vector src/Chat.java model.gguf
 *   java --add-modules jdk.incubator.vector src/Chat.java model.gguf --temp 0 --ctx 4096
 *   java --add-modules jdk.incubator.vector src/Chat.java model.gguf --seed 42 --systeem "Je bent een dichter."
 * </pre>
 *
 * <p>Drie dingen maken dit een echt gesprek in plaats van een reeks losse vragen:
 * <ol>
 *   <li><b>Het geheugen blijft staan.</b> De KV-cache loopt gewoon door over de
 *       beurten heen: per beurt worden alleen de nieuwe tokens gevoerd (jouw vraag
 *       plus de sjabloonranden), nooit het hele gesprek opnieuw. Het model onthoudt
 *       dus wat je eerder zei &mdash; tot de context vol is.</li>
 *   <li><b>Sampling.</b> Standaard temperatuur 0,8 met top-k 40 en top-p 0,9
 *       (dezelfde standaardwaarden als ollama). De gebruikte seed wordt altijd
 *       getoond; met {@code --seed} is een gesprek exact te herhalen, en met
 *       {@code --temp 0} is het volledig deterministisch.</li>
 *   <li><b>Streaming zonder halve tekens.</b> Tokens verschijnen zodra ze er zijn,
 *       maar de bytes worden gebufferd tot ze een compleet UTF-8-teken vormen
 *       ({@link Utf8Stroom}) &mdash; de valkuil uit hoofdstuk 13, hier definitief
 *       opgelost.</li>
 * </ol>
 *
 * <p>In het gesprek: {@code /nieuw} begint opnieuw, {@code /stop} (of Ctrl-D) stopt.
 */
public final class Chat {

    static final String GEBRUIK = "gebruik: java --add-modules jdk.incubator.vector src/Chat.java "
            + "<model.gguf> [--temp T] [--top-k K] [--top-p P] [--seed N] [--ctx N] [--systeem \"...\"]";

    public static void main(String[] args) throws Exception {
        if (args.length < 1) { System.err.println(GEBRUIK); System.exit(2); }
        Path model = Path.of(args[0]);
        if (!Files.isReadable(model)) {
            System.err.println("kan het modelbestand niet lezen: " + model.toAbsolutePath());
            System.exit(1);
        }

        float temp = Sampler.STANDAARD_TEMPERATUUR;
        int topK = Sampler.STANDAARD_TOP_K;
        float topP = Sampler.STANDAARD_TOP_P;
        long seed = System.nanoTime() & 0xFFFF;              // kort en toonbaar; --seed voor herhaling
        int ctx = 2048;
        String systeem = null;
        for (int i = 1; i < args.length; i++) {
            switch (args[i]) {
                case "--temp"    -> temp = Float.parseFloat(eis(args, ++i, "--temp"));
                case "--top-k"   -> topK = Integer.parseInt(eis(args, ++i, "--top-k"));
                case "--top-p"   -> topP = Float.parseFloat(eis(args, ++i, "--top-p"));
                case "--seed"    -> seed = Long.parseLong(eis(args, ++i, "--seed"));
                case "--ctx"     -> ctx = Integer.parseInt(eis(args, ++i, "--ctx"));
                case "--systeem" -> systeem = eis(args, ++i, "--systeem");
                default -> { System.err.println("onbekende optie: " + args[i]); System.err.println(GEBRUIK); System.exit(2); }
            }
        }

        try (Gguf g = Gguf.open(model)) {
            Tokenizer tk = new Tokenizer(g);
            ChatTemplate sjabloon = ChatTemplate.van(g);
            Motor m = Motor.open(g, ctx);
            Sampler sampler = new Sampler(temp, topK, topP, seed);

            System.out.printf("model: %s  |  context: %d tokens  |  temp %.2f, top-k %d, top-p %.2f, seed %d%n",
                    g.has("general.name") ? g.getString("general.name") : model.getFileName(),
                    m.ctx(), temp, topK, topP, seed);
            System.out.println("typ je vraag; /nieuw begint opnieuw, /stop stopt.");

            int[] open = tk.encode(sjabloon.opening(systeem));
            if (open.length + 64 > m.ctx()) {
                System.err.printf("de context (%d tokens) is te klein voor de systeemtekst (%d tokens) "
                        + "plus een vraag; kies --ctx groter of een kortere --systeem%n", m.ctx(), open.length);
                System.exit(1);
            }
            float[] logits = voer(m, open);

            BufferedReader in = new BufferedReader(new InputStreamReader(System.in, StandardCharsets.UTF_8));
            while (true) {
                System.out.print("\njij>  ");
                String regel = in.readLine();
                if (regel == null || regel.trim().equals("/stop")) break;
                regel = regel.trim();
                if (regel.isEmpty()) continue;
                if (regel.equals("/nieuw")) {
                    m.reset();
                    logits = voer(m, open);
                    System.out.println("(nieuw gesprek)");
                    continue;
                }

                // de sjabloonranden mét, de getypte tekst zónder speciale tokens: wie
                // <|im_end|> intypt, stuurt die zeven tekens — geen vervalste beurtwissel
                int[] delta = samengevoegd(tk.encode(sjabloon.vraagVoor()),
                        tk.encode(regel, false), tk.encode(sjabloon.vraagNa()));
                if (m.position() + delta.length + 32 > m.ctx()) {
                    System.out.printf("(het gesprek zit vol: %d van %d tokens — typ /nieuw, of start met --ctx groter)%n",
                            m.position(), m.ctx());
                    continue;
                }
                for (int id : delta) logits = m.forward(id);

                System.out.print("model> ");
                Utf8Stroom uit = new Utf8Stroom();
                long t0 = System.nanoTime();
                int aantal = 0;
                int volgende = sampler.kies(logits);
                boolean afgekapt = false;
                // qwen3.5-denkmodellen openen elk antwoord met <think>…</think>; die twee
                // tokens horen bij het antwoord en stromen mee (bij qwen2 bestaan ze niet)
                int denkOpen = tk.idOf("<think>"), denkDicht = tk.idOf("</think>");
                while (true) {
                    boolean denk = volgende == denkOpen || volgende == denkDicht;
                    if (!denk && (volgende == tk.eosId || tk.isSpecial(volgende))) break;
                    System.out.print(uit.voeg(tk.tokenBytes(volgende)));
                    System.out.flush();
                    aantal++;
                    if (m.position() + 3 > m.ctx()) {          // houd 2 slots voor <|im_end|>\n
                        afgekapt = true;
                        break;
                    }
                    logits = m.forward(volgende);
                    volgende = sampler.kies(logits);
                }
                System.out.print(uit.rest());
                // de geschiedenis sluit in álle gevallen met <|im_end|>\n — zoals llama.cpp
                // vervangt dat ook een afwijkend stop-token (bv. <|endoftext|>), en een
                // afgekapt antwoord krijgt zo alsnog een geldige beurtwissel
                logits = voerBinnen(m, tk, logits, "<|im_end|>\n");
                double s = (System.nanoTime() - t0) / 1e9;
                System.out.printf("%n(%d tokens, %.1f per seconde; gesprek: %d/%d)%n",
                        aantal, aantal / Math.max(s, 1e-9), m.position(), m.ctx());
                if (afgekapt) System.out.println("(het antwoord werd afgekapt: de context is vol — typ /nieuw)");
            }
            System.out.println("tot ziens.");
        }
    }

    private static float[] voer(Motor m, int[] ids) {
        float[] logits = null;
        for (int id : ids) logits = m.forward(id);
        return logits;
    }

    /** Voert tokens zolang de context het toelaat; wat niet past valt weg (de beurtwachter vangt dat op). */
    private static float[] voerBinnen(Motor m, Tokenizer tk, float[] logits, String tekst) {
        for (int id : tk.encode(tekst)) {
            if (m.position() >= m.ctx()) break;
            logits = m.forward(id);
        }
        return logits;
    }

    private static int[] samengevoegd(int[]... delen) {
        int n = 0;
        for (int[] d : delen) n += d.length;
        int[] uit = new int[n];
        int i = 0;
        for (int[] d : delen) { System.arraycopy(d, 0, uit, i, d.length); i += d.length; }
        return uit;
    }

    private static String eis(String[] args, int i, String vlag) {
        if (i >= args.length) { System.err.println("optie " + vlag + " heeft een waarde nodig"); System.exit(2); }
        return args[i];
    }

    /**
     * Buffert bytes tot ze een compleet UTF-8-teken vormen. E&eacute;n token kan een
     * halve tekencode bevatten (hoofdstuk 13); wie per token afdrukt zonder buffer,
     * drukt vervangtekens af. Deze stroom geeft telkens het langste decodeerbare
     * voorstuk terug en houdt de rest vast tot het vervolg er is.
     */
    static final class Utf8Stroom {
        private byte[] buf = new byte[64];
        private int n;

        String voeg(byte[] bytes) {
            if (n + bytes.length > buf.length) {
                buf = java.util.Arrays.copyOf(buf, Math.max(buf.length * 2, n + bytes.length));
            }
            System.arraycopy(bytes, 0, buf, n, bytes.length);
            n += bytes.length;

            int compleet = n;
            int i = n - 1;
            int staart = 0;
            // een teken telt hoogstens 3 vervolgbytes; verder terugkijken hoeft nooit
            while (i >= 0 && staart < 3 && (buf[i] & 0xC0) == 0x80) { i--; staart++; }
            if (i < 0) {
                compleet = 0;                                 // alleen vervolgbytes: vasthouden
            } else if ((buf[i] & 0xC0) == 0x80) {
                compleet = n;                                 // 4+ vervolgbytes op rij: kapot, doorlaten
            } else {
                int kop = buf[i] & 0xFF;
                int nodig = kop >= 0xF0 ? 4 : kop >= 0xE0 ? 3 : kop >= 0xC0 ? 2 : 1;
                if (staart + 1 < nodig) compleet = i;         // de laatste reeks is nog niet af
            }
            String uitvoer = new String(buf, 0, compleet, StandardCharsets.UTF_8);
            System.arraycopy(buf, compleet, buf, 0, n - compleet);
            n -= compleet;
            return uitvoer;
        }

        /** Wat er nog vastzit — aan het einde van een antwoord, desnoods met vervangteken. */
        String rest() {
            String s = new String(buf, 0, n, StandardCharsets.UTF_8);
            n = 0;
            return s;
        }
    }
}
