#!/usr/bin/env python3
"""
./gen_queries.py frequency_dictionary_en_82_765.txt queries.txt --corpus pg11.txt

Dictionary: https://github.com/wolfgarbe/SymSpell/blob/master/SymSpell/frequency_dictionary_en_82_765.txt
Corpus (Alice in Wonderland): https://www.gutenberg.org/cache/epub/11/pg11.txt

Sections:
@HITS    5000 dictionary words
@TYPOS   same words with 1 or 2 random edits
@MISSES  5000 random strings of 7 to 12 letters
@PHRASES input<TAB>original words. With a corpus, 200 sentences of 4 to 9 words from Alice in Wonderland.
         four variants: spaces removed, spaces removed plus a typo, random add/remove spaces, random add/remove spaces plus a typo
@TEXTS   error percent<TAB>misspelled document<TAB>original, 1200 frequent words each
"""

import argparse
import random
import re

QUERIES = 5000
PHRASES = 200
TEXT_WORDS = 1200
TEXT_ERROR_PERCENTS = (5, 10, 20)
COMMON = 20_000
SEED = 0x3C38E88DF7E69E67

def typo(rng, word, edits):
    s = list(word)
    for _ in range(edits):
        pos = rng.randrange(len(s))
        letter = chr(ord("a") + rng.randrange(26))
        op = rng.randrange(4)
        if op == 0:
            if len(s) > 1:
                del s[pos]
        elif op == 1:
            s.insert(pos, letter)
        elif op == 2:
            s[pos] = letter
        elif pos + 1 < len(s):
            s[pos], s[pos + 1] = s[pos + 1], s[pos]

    return "".join(s)


def sentences(paths, known, rng):
    out, seen = [], set()
    for path in paths:
        text = open(path, encoding="utf-8").read()
        start, end = text.find("*** START"), text.find("*** END")

        if 0 <= start < end:
            text = text[start:end]

        for chunk in re.split(r"[.!?;:]", text.replace("\r", "").replace("\n", " ")):
            ws = re.findall(r"[a-z]+", chunk.lower())
            key = " ".join(ws)
            if 4 <= len(ws) <= 9 and all(w in known for w in ws) and key not in seen:
                seen.add(key)
                out.append(ws)

    rng.shuffle(out)
    return out


def with_typo(rng, ws):
    ws = list(ws)
    i = rng.randrange(len(ws))
    ws[i] = typo(rng, ws[i], 1)
    return ws


def noisy(rng, ws):
    ws = list(ws)
    if rng.random() < 0.5:
        i = rng.randrange(len(ws))
        if len(ws[i]) > 2:
            cut = rng.randrange(1, len(ws[i]))
            ws[i] = ws[i][:cut] + " " + ws[i][cut:]

    out = ws[0]
    for w in ws[1:]:
        out += (" " if rng.random() < 0.5 else "") + w

    return out


def main():
    ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
    ap.add_argument("dictionary", help="one `word count` per line")
    ap.add_argument("out")
    ap.add_argument("--seed", type=lambda s: int(s, 0), default=SEED)
    ap.add_argument("--corpus", action="append", default=[], help="plain text file to take sentences from")
    args = ap.parse_args()

    with open(args.dictionary, encoding="utf-8-sig") as f:
        words = [fields[0] for fields in map(str.split, f) if len(fields) == 2]
    rng = random.Random(args.seed)

    hits = [rng.choice(words) for _ in range(QUERIES)]
    typos = [typo(rng, rng.choice(hits), 1 + rng.randrange(2)) for _ in range(QUERIES)]
    misses = ["".join(chr(ord("a") + rng.randrange(26)) for _ in range(7 + rng.randrange(6))) for _ in range(QUERIES)]

    common = words[:COMMON]
    phrases = []
    if args.corpus:
        for i, parts in enumerate(sentences(args.corpus, set(words), rng)[:PHRASES]):
            typed = with_typo(rng, parts) if i % 2 == 1 else parts
            phrases.append((noisy(rng, typed) if i % 4 >= 2 else "".join(typed), parts))
    else:
        for i in range(PHRASES):
            parts = [rng.choice(common) for _ in range(3 + rng.randrange(3))]
            at = rng.randrange(len(parts)) if i % 2 == 0 else len(parts)
            phrases.append(("".join(typo(rng, w, 1) if k == at else w for k, w in enumerate(parts)), parts))

    texts = []
    for percent in TEXT_ERROR_PERCENTS:
        original = [rng.choice(common) for _ in range(TEXT_WORDS)]
        typed = [typo(rng, w, 1) if rng.randrange(100) < percent else w for w in original]
        texts.append((percent, typed, original))

    with open(args.out, "w", encoding="utf-8") as f:
        f.write(f"# symspell bench queries, seed 0x{args.seed:X}\n")

        for name, entries in (("HITS", hits), ("TYPOS", typos), ("MISSES", misses)):
            f.write(f"@{name}\n")
            f.writelines(f"{it}\n" for it in entries)

        f.write("@PHRASES\n")
        f.writelines(f"{glued}\t{' '.join(parts)}\n" for glued, parts in phrases)

        f.write("@TEXTS\n")
        f.writelines(f"{p}\t{' '.join(t)}\t{' '.join(o)}\n" for p, t, o in texts)

    print(f"{args.out}: {len(words)} words, "
          f"{QUERIES} per set, "
          f"{len(phrases)} phrases, "
          f"{len(texts)} texts of "
          f"{TEXT_WORDS} words, "
          f"seed 0x{args.seed:X}")


if __name__ == "__main__":
    main()
