"""A complete word-token count model. Core example: Python 3.10+, standard library.

Run this file to print every training row and save the result JSON. Add --chart
to also draw the chart (requires matplotlib and the course_utils.py helper).
Each supplied line is a separate sequence; blank lines are ignored. To model
one long passage, pass it as a single list entry, including any line breaks.
"""
import argparse
import json
import random
import unicodedata
from collections import Counter, defaultdict
from pathlib import Path

BOS = "<BOS>"
EOS = "<EOS>"
CORPUS = [
    "the bug is urgent",
    "the bug is urgent",
    "the bug is reproducible",
    "the bug is broken",
    "the build is broken",
]


def tokenize(text):
    """Lowercase word tokens; retain letters, digits and internal apostrophes."""
    text = unicodedata.normalize("NFC", text).lower().replace("’", "'")
    words, current = [], []
    for index, char in enumerate(text):
        kind = unicodedata.category(char)[0]
        if kind in "LN" or (kind == "M" and current and current[-1] != "'"):
            current.append(char)
        elif char == "'" and current and current[-1] != "'" and index + 1 < len(text) and unicodedata.category(text[index + 1])[0] in "LN":
            current.append(char)
        elif current:
            words.append("".join(current))
            current = []
    if current:
        words.append("".join(current))
    return words


def train_ngram(lines, order=3):
    if order not in (1, 2, 3):
        raise ValueError("Choose order 1, 2, or 3")
    counts = defaultdict(Counter)
    rows = []
    for line in lines:
        words = tokenize(line)
        if not words:
            continue
        tokens = [BOS] * (order - 1) + words + [EOS]
        for index in range(order - 1, len(tokens)):
            context = tuple(tokens[index - order + 1:index])
            label = tokens[index]
            counts[context][label] += 1
            rows.append((context, label))
    return counts, rows


def context_for(history, order=3):
    if order not in (1, 2, 3):
        raise ValueError("Choose order 1, 2, or 3")
    width = order - 1
    return tuple(([BOS] * width + list(history))[-width:]) if width else ()


def distribution(counts, context):
    """Exact-context lookup. An unseen context has no estimate, returned as {}."""
    bucket = counts.get(tuple(context), {})
    total = sum(bucket.values())
    return {token: count / total for token, count in sorted(bucket.items())} if total else {}


def generate(counts, prompt="", order=3, steps=30, seed=7, sample=False):
    """Reuse the fitted table. Training itself always uses the supplied text."""
    history = tokenize(prompt)
    rng = random.Random(seed)
    trace = []
    for _ in range(steps):
        context = context_for(history, order)
        probabilities = distribution(counts, context)
        if not probabilities:
            return history, trace, "unseen context: no estimate"
        ranked = sorted(probabilities, key=lambda token: (-probabilities[token], token))
        target = rng.choices(ranked, weights=[probabilities[t] for t in ranked])[0] if sample else ranked[0]
        trace.append({"context": list(context), "target": target, "probability": probabilities[target]})
        if target == EOS:
            return history, trace, "end marker selected"
        history.append(target)
    return history, trace, "step limit reached"


def example_result():
    counts, rows = train_ngram(CORPUS)
    ordinary = sorted({token for line in CORPUS for token in tokenize(line)})
    context = ("bug", "is")
    probabilities = distribution(counts, context)
    history, trace, stop = generate(counts)
    return {
        "corpus": CORPUS,
        "boundary_convention": {"start": BOS, "end": EOS, "start_markers": 2, "end_markers": 1},
        "ordinary_vocabulary": ordinary,
        "input_vocabulary": ordinary + [BOS, EOS],
        "output_vocabulary": ordinary + [EOS],
        "training_rows": [[list(ctx), label] for ctx, label in rows],
        "query_context": list(context),
        "continuation_counts": dict(counts[context]),
        "continuation_probabilities": probabilities,
        "top_prediction": max(probabilities, key=probabilities.get),
        "unseen_context_distribution": distribution(counts, ("is", "the")),
        "generation": {"tokens": history, "trace": trace, "stop": stop},
    }


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--chart", action="store_true", help="Also generate the matplotlib chart")
    args = parser.parse_args()
    result = example_result()
    root = Path(__file__).resolve().parent
    (root / "results").mkdir(exist_ok=True)
    (root / "results/04_ngram_language_model.json").write_text(json.dumps(result, indent=2) + "\n")
    print("TRIGRAM: TWO START MARKERS, ONE END TARGET PER SEQUENCE")
    for context, target in result["training_rows"]:
        print(f"{context} -> {target}")
    print(f"training examples = {len(result['training_rows'])}")
    print("P(next | bug is) =", result["continuation_probabilities"])
    print("GENERATION TRACE (most likely token; fixed tie-break)")
    for step in result["generation"]["trace"]:
        print(f"{step['context']} -> {step['target']} (p={step['probability']:.3f})")
    print("Generated:", " ".join(result["generation"]["tokens"]), "—", result["generation"]["stop"])
    if args.chart:
        from course_utils import BLUE, CYAN, GREEN, save_chart, setup_chart
        import matplotlib.pyplot as plt
        setup_chart()
        ranked = sorted(result["continuation_probabilities"].items(), key=lambda item: (-item[1], item[0]))
        fig, ax = plt.subplots(figsize=(10.5, 5.4))
        bars = ax.bar([t for t, _ in ranked], [p for _, p in ranked], color=[BLUE, CYAN, GREEN], width=0.62)
        ax.set(title='What followed “bug is” in the training text?', xlabel="candidate next token", ylabel="matching occurrences / context total", ylim=(0, .62))
        ax.grid(axis="y")
        total = sum(result["continuation_counts"].values())
        for bar, (token, probability) in zip(bars, ranked):
            ax.text(bar.get_x() + bar.get_width()/2, probability+.025, f"{result['continuation_counts'][token]}/{total} = {probability:.2f}", ha="center", fontweight="bold")
        print("chart =", save_chart(fig, "04_ngram_language_model.png"))


if __name__ == "__main__":
    main()
