import java.nio.charset.StandardCharsets;
import java.nio.file.Files;
import java.nio.file.Path;
import java.util.Arrays;
import java.util.List;

/**
 * Toont hoe een stuk tekst in tokens uiteenvalt. Klein gereedschap, maar onmisbaar zodra er
 * iets misgaat: bij een model dat wartaal uitkraamt is de eerste vraag altijd of de prompt wel
 * juist opgedeeld is.
 *
 * <pre>
 *   java src/Tok.java model.gguf "Hallo wereld"
 *   java src/Tok.java model.gguf --file tekst.txt
 *   java src/Tok.java model.gguf --count "alleen het aantal"
 *   java src/Tok.java model.gguf --chunks "toon ook de voorsplitsing"
 *   java src/Tok.java model.gguf --decode 9707 1879
 * </pre>
 */
public final class Tok {

    public static void main(String[] args) throws Exception {
        if (args.length < 2) {
            System.err.println("gebruik: java src/Tok.java <model.gguf> [--count|--chunks|--file|--decode] <tekst...>");
            System.exit(2);
        }
        try (Gguf g = Gguf.open(Path.of(args[0]))) {
            Tokenizer tk = new Tokenizer(g);

            // vlaggen mogen gecombineerd worden en staan altijd vooraan
            boolean count = false, chunks = false, file = false, decode = false;
            int from = 1;
            while (from < args.length && args[from].startsWith("--")) {
                switch (args[from]) {
                    case "--count"  -> count = true;
                    case "--chunks" -> chunks = true;
                    case "--file"   -> file = true;
                    case "--decode" -> decode = true;
                    default -> {
                        System.err.println("onbekende optie: " + args[from]);
                        System.exit(2);
                    }
                }
                from++;
            }
            if (from >= args.length) {
                System.err.println("geen tekst opgegeven");
                System.exit(2);
            }

            if (decode) {
                int[] ids = new int[args.length - from];
                for (int i = from; i < args.length; i++) ids[i - from] = Integer.parseInt(args[i]);
                System.out.println(tk.decode(ids));
                return;
            }

            String text = file
                    ? Files.readString(Path.of(args[from]), StandardCharsets.UTF_8)
                    : String.join(" ", Arrays.asList(args).subList(from, args.length));

            int[] ids = tk.encode(text);

            if (count) {
                System.out.println(ids.length);
                return;
            }

            if (chunks) {
                List<String> parts = tk.chunks(text);
                System.out.println("voorsplitsing (" + parts.size() + " brokken):");
                for (String c : parts) System.out.println("  " + show(c));
                System.out.println();
            }

            System.out.printf("%,d tekens, %,d bytes, %,d tokens  (%.2f tekens per token)%n",
                    text.length(), text.getBytes(StandardCharsets.UTF_8).length, ids.length,
                    ids.length == 0 ? 0 : (double) text.length() / ids.length);
            System.out.println();
            System.out.printf("  %6s  %8s  %-24s %s%n", "#", "id", "byte-level", "als tekst");
            String[] pieces = tk.pieces(ids);
            for (int i = 0; i < ids.length; i++) {
                String asText = tk.decode(ids, i, i + 1);
                System.out.printf("  %6d  %8d  %-24s %s%s%n", i, ids[i], show(pieces[i]), show(asText),
                        tk.isSpecial(ids[i]) ? "   <- speciaal token" : "");
            }
            System.out.println();
            System.out.println("nummers: " + Arrays.toString(ids));

            String back = tk.decode(ids);
            System.out.println(back.equals(text)
                    ? "heen en terug: ongeschonden"
                    : "heen en terug: AFWIJKING -> " + show(back));
        }
    }

    static String show(String s) {
        String t = s.replace("\\", "\\\\").replace("\n", "\\n").replace("\t", "\\t").replace("\r", "\\r");
        if (t.length() > 22) t = t.substring(0, 19) + "...";
        return t;
    }
}
