#!/usr/bin/env python3.13
"""Write down what Anki itself does, as reference numbers for Loudwalk's scheduler.

Loudwalk promises an Anki user the behaviour they already know: the same learning steps,
the same daily limits, the same order of cards, the same interval under each button. The
only honest proof of "the same" is Anki's own answer to the same question, so this script
asks the Anki installed on this Mac and records what it says in data/anki_golden.json.
engine/Tests/FSRSEngineTests/AnkiParityTests.swift then requires Loudwalk to give every
one of those answers (`cd engine && swift test`).

It works in throwaway collections under a temporary directory and never opens a profile.
Nothing here is part of any Loudwalk build, and no Anki code is linked into the app: Anki
is AGPL, so it is only ever run, as a black box, and its outputs are plain numbers.

    PYTHONPATH=/Applications/Anki.app/Contents/Resources/app_packages \
        python3.13 tools/anki_golden.py

Recorded against Anki 26.08.1 (the version the output names). Re-run after upgrading Anki
and look at what changed before accepting it.
"""

import json
import math
import os
import random
import sys
import tempfile
import time

from anki.buildinfo import version as anki_version
from anki.collection import Collection
from anki.scheduler_pb2 import CardAnswer

try:
    from anki.cards import FSRSMemoryState
except ImportError:  # older layouts
    from anki.cards_pb2 import FsrsMemoryState as FSRSMemoryState

REPO = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
OUT = os.path.join(REPO, "data", "anki_golden.json")

# HdA's own fitted parameters (the preset most new cards go through) and FSRS-6's defaults.
with open(os.path.join(REPO, "data", "bundle.json")) as f:
    HDA = json.load(f)["schedulers"]["HdA"]["parameters"]
DEFAULTS = [
    0.212, 1.2931, 2.3065, 8.2956, 6.4133, 0.8334, 3.0194, 0.001, 1.8722, 0.1666, 0.796,
    1.4835, 0.0614, 0.2629, 1.6483, 0.6014, 1.8729, 0.5425, 0.0912, 0.0658, 0.1542,
]
STEPS, RELEARN, MAX_IVL, RETENTION = [1.0, 10.0], [10.0], 36500, 0.9
RATINGS = [CardAnswer.AGAIN, CardAnswer.HARD, CardAnswer.GOOD, CardAnswer.EASY]
rng = random.Random(20260918)


def collection(tmp, name, params, balance, days_back=0):
    path = os.path.join(tmp, name + ".anki2")
    col = Collection(path)
    if days_back:
        # An old collection, so "today" is a day number with a past behind it. The
        # backend caches the creation date, hence the reopen.
        col.db.execute("update col set crt = crt - ?", days_back * 86400)
        col.close()
        col = Collection(path)
        assert col.sched.today == days_back, col.sched.today
    col.set_config("fsrs", True)
    col.set_config("loadBalancerEnabled", balance)
    did = col.decks.id("Deck")
    conf = col.decks.config_dict_for_deck_id(did)
    conf["new"]["delays"] = STEPS
    conf["lapse"]["delays"] = RELEARN
    conf["fsrsParams6"] = params
    conf["desiredRetention"] = RETENTION
    conf["rev"]["maxIvl"] = MAX_IVL
    conf["new"]["perDay"] = 20
    conf["rev"]["perDay"] = 200
    col.decks.update_config(conf)
    return col, did, conf


def add_cards(col, did, count, reversed_too=False):
    model = col.models.by_name("Basic (and reversed card)" if reversed_too else "Basic")
    ids = []
    for i in range(count):
        note = col.new_note(model)
        note["Front"], note["Back"] = f"q{i}-{rng.random()}", f"a{i}"
        col.add_note(note, did)
        ids += [c.id for c in note.cards()]
    return ids


def put(col, cid, **fields):
    card = col.get_card(cid)
    memory = fields.pop("memory", None)
    for key, value in fields.items():
        setattr(card, key, value)
    if memory is not None:
        card.memory_state = FSRSMemoryState(stability=memory[0], difficulty=memory[1])
    col.update_card(card, skip_undo_entry=True)
    return col.get_card(cid)


def outcome(state):
    """One of the four states under the buttons, as plain numbers."""
    normal = state.normal
    kind = normal.WhichOneof("kind")
    if kind == "learning":
        l = normal.learning
        return {"kind": "learning", "remaining": l.remaining_steps, "secs": l.scheduled_secs,
                "s": l.memory_state.stability, "d": l.memory_state.difficulty}
    if kind == "review":
        r = normal.review
        return {"kind": "review", "days": r.scheduled_days,
                "s": r.memory_state.stability, "d": r.memory_state.difficulty}
    if kind == "relearning":
        r = normal.relearning
        return {"kind": "relearning", "remaining": r.learning.remaining_steps,
                "secs": r.learning.scheduled_secs, "days": r.review.scheduled_days,
                "s": r.learning.memory_state.stability,
                "d": r.learning.memory_state.difficulty}
    raise ValueError(kind)


def elapsed_days(col, last_review):
    return max(0, col.sched.day_cutoff - last_review) // 86400


def scenario(col, cid, params):
    """Everything the engine needs to recompute Anki's four answers for this card."""
    card = col.get_card(cid)
    states = col._backend.get_scheduling_states(cid)
    memory = card.memory_state
    has_memory = card.type != 0 and memory is not None and memory.stability > 0
    kinds = {0: "new", 1: "learning", 2: "review", 3: "relearning"}
    return {
        "cid": cid, "reps": card.reps, "kind": kinds[card.type],
        "remaining": card.left % 1000, "scheduled_days": card.ivl,
        "s": memory.stability if has_memory else None,
        "d": memory.difficulty if has_memory else None,
        "elapsed_days": elapsed_days(col, card.last_review_time) if card.last_review_time else 0,
        "secs_until_rollover": col.sched.day_cutoff - int(time.time()),
        "params": params,
        "expect": [outcome(getattr(states, n)) for n in ("again", "hard", "good", "easy")],
    }


def fuzz_section(tmp):
    """The random draw itself: one seed, one interval, the day Anki picks."""
    col, did, _ = collection(tmp, "fuzz", HDA, balance=False)
    cases = []
    for cid in add_cards(col, did, 3):
        for reps in list(range(1, 120)) + [rng.randrange(120, 60000) for _ in range(40)]:
            col.db.execute("update cards set reps = ? where id = ?", reps, cid)
            for interval in (30000, 3, 7, 17, 37, 100):
                delta = col._backend.fuzz_delta(card_id=cid, interval=interval)
                # Anki seeds the fuzz with card id + reps; for rescheduling, reps - 1.
                cases.append({"seed": cid + reps - 1, "interval": interval,
                              "max": MAX_IVL, "days": interval + delta})
    col.close()
    return cases


def learning_fuzz_section(tmp):
    """The extra seconds Anki adds to a learning step, seen by answering real cards."""
    col, did, _ = collection(tmp, "learnfuzz", HDA, balance=False)
    cases = []
    ids = add_cards(col, did, 150)
    for i, cid in enumerate(ids):
        for rating in ([RATINGS[i % 3]] + ([RATINGS[(i // 3) % 3]] if i % 2 else [])):
            card = col.get_card(cid)
            reps, states = card.reps, col._backend.get_scheduling_states(cid)
            chosen = outcome(getattr(states, ("again", "hard", "good", "easy")[rating]))
            if chosen["kind"] != "learning":
                continue
            card.start_timer()
            before = int(time.time())
            col.sched.answer_card(col.sched.build_answer(card=card, states=states, rating=rating))
            after = int(time.time())
            due = col.get_card(cid).due
            cases.append({"seed": cid + reps, "secs": chosen["secs"],
                          "low": due - after, "high": due - before})
    col.close()
    return cases


def transitions_section(tmp):
    """The four answers for cards in every state, with plain fuzz."""
    out = []
    for label, params in (("hda", HDA), ("default", DEFAULTS)):
        col, did, _ = collection(tmp, "states-" + label, params, balance=False)
        now = int(time.time())
        cutoff = col.sched.day_cutoff
        ids = iter(add_cards(col, did, 400))
        # New cards.
        for _ in range(15):
            cid = next(ids)
            put(col, cid, reps=rng.randrange(0, 5))
            out.append(scenario(col, cid, params))
        # Learning, both steps, same day and from earlier days.
        for _ in range(80):
            cid = next(ids)
            left = rng.choice([1, 2])
            same_day = rng.random() < 0.6
            last = now - rng.randrange(30, 3000) if same_day else \
                cutoff - 86400 * rng.randrange(1, 6) - rng.randrange(60, 80000)
            put(col, cid, type=1, queue=1, left=left, due=now, reps=rng.randrange(1, 12),
                memory=(rng.uniform(0.05, 12), rng.uniform(1, 10)), last_review_time=last)
            out.append(scenario(col, cid, params))
        # Review cards, on time, late and early, short and long.
        for _ in range(220):
            cid = next(ids)
            ivl = rng.choice([1, 2, 3, 4, 6, 9, 14, 21, 35, 60, 90, 150, 300, 800])
            elapsed = max(0, ivl + rng.randrange(-ivl // 2 - 1, ivl + 3))
            stability = max(0.1, ivl * rng.uniform(0.3, 2.5))
            last = cutoff - 86400 * elapsed - rng.randrange(60, 80000)
            put(col, cid, type=2, queue=2, ivl=ivl, due=col.sched.today, reps=rng.randrange(1, 60),
                memory=(stability, rng.uniform(1, 10)), last_review_time=last)
            out.append(scenario(col, cid, params))
        # Relearning after a lapse.
        for _ in range(60):
            cid = next(ids)
            same_day = rng.random() < 0.7
            last = now - rng.randrange(30, 3000) if same_day else \
                cutoff - 86400 * rng.randrange(1, 4) - rng.randrange(60, 80000)
            put(col, cid, type=3, queue=1, left=1, due=now, ivl=rng.randrange(1, 200),
                reps=rng.randrange(3, 80), memory=(rng.uniform(0.05, 30), rng.uniform(1, 10)),
                last_review_time=last)
            out.append(scenario(col, cid, params))
        col.close()
    return out


def balance_section(tmp):
    """The load balancer: the same questions with a known pile of cards on each day."""
    col, did, _ = collection(tmp, "balance", HDA, balance=True, days_back=200)
    today = col.sched.today
    cutoff = col.sched.day_cutoff
    now = int(time.time())
    background = add_cards(col, did, 900)
    pile = iter(background)
    for day in range(1, 99):
        for _ in range(rng.choice([0, 0, 1, 2, 3, 5, 8, 12, 20])):
            cid = next(pile, None)
            if cid is None:
                break
            put(col, cid, type=2, queue=2, ivl=day, due=today + day, reps=3,
                memory=(day, 5.0), last_review_time=now - 86400 * 3)
    # Whatever is left over is suspended out of the way at a far day.
    for cid in pile:
        put(col, cid, type=2, queue=-1, ivl=500, due=today + 500, reps=3,
            memory=(500, 5.0), last_review_time=now - 86400 * 3)
    tests = add_cards(col, did, 160)
    for i, cid in enumerate(tests):
        if i < 40:
            left = rng.choice([1, 2])
            put(col, cid, type=1, queue=1, left=left, due=now, reps=rng.randrange(1, 9),
                memory=(rng.uniform(0.5, 25), rng.uniform(1, 10)),
                last_review_time=now - rng.randrange(30, 3000))
        else:
            ivl = rng.choice([2, 3, 5, 8, 13, 21, 34, 55, 80])
            elapsed = max(1, ivl + rng.randrange(-1, 4))
            put(col, cid, type=2, queue=2, ivl=ivl, due=today, reps=rng.randrange(1, 40),
                memory=(max(0.5, ivl * rng.uniform(0.6, 2.2)), rng.uniform(1, 10)),
                last_review_time=cutoff - 86400 * elapsed - rng.randrange(60, 80000))
    col.decks.select(did)
    col.sched.get_queued_cards()  # builds the queues, and with them the balancer
    counts = dict(col.db.all(
        "select due - ?, count() from cards where due >= ? and due < ? group by due",
        today, today, today + 99))
    cases = [scenario(col, cid, HDA) for cid in tests]
    col.close()
    return {"load": {str(k): v for k, v in counts.items()}, "cases": cases}


def queue_section(tmp):
    """Which card Anki shows, in what order, under which limits."""
    out = []
    setups = [
        # (reviews due, interday learning, new notes, rev/day, new/day, done today new/rev)
        (30, 2, 14, 200, 20, 0, 0),
        (50, 0, 10, 45, 20, 0, 0),
        (50, 0, 12, 60, 20, 0, 0),
        (30, 3, 12, 60, 20, 5, 10),
        (8, 0, 25, 200, 20, 0, 0),
        (120, 4, 30, 200, 20, 0, 0),
        (0, 0, 9, 200, 5, 2, 0),
        (40, 0, 0, 30, 20, 0, 0),
    ]
    for n, (reviews, day_learning, new_notes, rev_limit, new_limit, done_new, done_rev) in \
            enumerate(setups):
        col, did, conf = collection(tmp, f"queue{n}", HDA, balance=False, days_back=300)
        conf["rev"]["perDay"], conf["new"]["perDay"] = rev_limit, new_limit
        col.decks.update_config(conf)
        today = col.sched.today
        now = int(time.time())
        if done_new or done_rev:
            col._backend.update_stats(deck_id=did, new_delta=done_new, review_delta=done_rev,
                                      millisecond_delta=0)
        for cid in add_cards(col, did, reviews):
            ivl = rng.randrange(1, 60)
            put(col, cid, type=2, queue=2, ivl=ivl, due=today - rng.randrange(0, 12),
                reps=4, memory=(ivl, 5.0), last_review_time=now - 86400 * ivl)
        for cid in add_cards(col, did, day_learning):
            put(col, cid, type=1, queue=3, left=1, due=today - rng.randrange(0, 3), reps=2,
                memory=(1.0, 5.0), last_review_time=now - 86400 * 2)
        for cid in add_cards(col, did, 4):  # intraday learning: due, soon, and later today
            put(col, cid, type=1, queue=1, left=rng.choice([1, 2]),
                due=now + rng.choice([-400, -60, 300, 900, 5000]), reps=rng.randrange(0, 3),
                memory=(0.5, 5.0), last_review_time=now - 600)
        new_ids = add_cards(col, did, new_notes, reversed_too=True)
        # Positions out of creation order, and a note sharing a position with another.
        positions = list(range(1, new_notes + 1))
        rng.shuffle(positions)
        for i, cid in enumerate(new_ids):
            put(col, cid, due=positions[i // 2] if i // 2 != 3 else positions[2])
        col.decks.select(did)
        queued = col.sched.get_queued_cards(fetch_limit=1000)
        kinds = {0: "new", 1: "learning", 2: "review"}
        cards = [
            {"cid": r[0], "nid": r[1], "ord": r[2], "mod": r[3], "type": r[4], "queue": r[5],
             "due": r[6], "reps": r[7]}
            for r in col.db.all("select id, nid, ord, mod, type, queue, due, reps from cards")
        ]
        out.append({
            "today": today, "cutoff": now, "next_day_at": col.sched.day_cutoff,
            "learn_ahead": col.get_config("collapseTime", 1200),
            "rev_limit": rev_limit, "new_limit": new_limit,
            "done_new": done_new, "done_rev": done_rev,
            "cards": cards,
            "order": [{"cid": q.card.id, "kind": kinds[q.queue]} for q in queued.cards],
            "counts": [queued.new_count, queued.learning_count, queued.review_count],
        })
        col.close()
    return out


def main():
    with tempfile.TemporaryDirectory(prefix="anki-golden-") as tmp:
        golden = {
            "anki": anki_version,
            "note": "Reference outputs recorded from Anki by tools/anki_golden.py. Plain "
                    "numbers; regenerate rather than edit.",
            "steps": [s * 60 for s in STEPS], "relearn": [s * 60 for s in RELEARN],
            "max_ivl": MAX_IVL, "retention": RETENTION,
            "fuzz": fuzz_section(tmp),
            "learning_fuzz": learning_fuzz_section(tmp),
            "transitions": transitions_section(tmp),
            "balance": balance_section(tmp),
            "queues": queue_section(tmp),
        }
    with open(OUT, "w") as f:
        json.dump(golden, f, separators=(",", ":"))
    print(f"wrote {OUT}: {len(golden['fuzz'])} fuzz, {len(golden['learning_fuzz'])} "
          f"learning fuzz, {len(golden['transitions'])} transitions, "
          f"{len(golden['balance']['cases'])} balanced, {len(golden['queues'])} queues "
          f"(Anki {anki_version})")


if __name__ == "__main__":
    sys.exit(main())
