#!/usr/bin/env python3
"""AB-100 drill engine: onboarding, session planning, answer capture, readiness reporting.

Usage:
  python drill.py setup --exam-date YYYY-MM-DD [--name "Your Name"] [--daily-target 20] [--force]
  python drill.py status
  python drill.py init
  python drill.py start [--count 20]
  python drill.py record --session N --qnum N --objective P.B.3 --correct 1 [--stem "..."]
  python drill.py finish --session N
  python drill.py report
  python drill.py history [--limit 10]
  python drill.py reset --confirm

All state lives beside this file: config.json (your exam date) and progress.db (your history).
"""

import argparse
import json
import os
import random
import sqlite3
import sys
from datetime import date, datetime

BASE = os.path.dirname(os.path.abspath(__file__))
DB = os.path.join(BASE, "progress.db")
BLUEPRINT = os.path.join(BASE, "blueprint.json")
CONFIG = os.path.join(BASE, "config.json")

ONBOARDING_MESSAGE = (
    "No exam date configured. Ask the user which date they are sitting AB-100, "
    'then run: python drill.py setup --exam-date YYYY-MM-DD [--name "Their Name"]'
)


def load_blueprint():
    with open(BLUEPRINT, "r", encoding="utf-8") as f:
        return json.load(f)


def load_config():
    if not os.path.exists(CONFIG):
        return None
    try:
        with open(CONFIG, "r", encoding="utf-8") as f:
            cfg = json.load(f)
    except (OSError, ValueError):
        return None
    if not cfg.get("exam_date"):
        return None
    return cfg


def save_config(cfg):
    with open(CONFIG, "w", encoding="utf-8") as f:
        json.dump(cfg, f, indent=2)
        f.write("\n")


def parse_exam_date(value):
    try:
        return date.fromisoformat(value.strip())
    except ValueError:
        raise SystemExit(
            json.dumps(
                {"ok": False, "error": f"Invalid date '{value}'. Use ISO format YYYY-MM-DD."},
                indent=2,
            )
        )


def require_config():
    """Return config, or emit an onboarding prompt and exit."""
    cfg = load_config()
    if cfg is None:
        print(json.dumps({"ok": False, "needs_onboarding": True, "action": ONBOARDING_MESSAGE}, indent=2))
        sys.exit(2)
    return cfg


def exam_date_from(cfg):
    return date.fromisoformat(cfg["exam_date"])


def days_until(cfg):
    return (exam_date_from(cfg) - date.today()).days


def conn():
    c = sqlite3.connect(DB)
    c.row_factory = sqlite3.Row
    return c


def ensure_schema():
    c = conn()
    c.executescript(
        """
        CREATE TABLE IF NOT EXISTS sessions (
            id INTEGER PRIMARY KEY AUTOINCREMENT,
            drill_date TEXT NOT NULL,
            started_at TEXT NOT NULL,
            completed_at TEXT,
            planned_count INTEGER NOT NULL DEFAULT 20
        );
        CREATE TABLE IF NOT EXISTS answers (
            id INTEGER PRIMARY KEY AUTOINCREMENT,
            session_id INTEGER NOT NULL,
            qnum INTEGER NOT NULL,
            objective_id TEXT NOT NULL,
            domain TEXT NOT NULL,
            correct INTEGER NOT NULL,
            stem TEXT,
            answered_at TEXT NOT NULL,
            UNIQUE (session_id, qnum),
            FOREIGN KEY (session_id) REFERENCES sessions (id)
        );
        CREATE INDEX IF NOT EXISTS idx_answers_obj ON answers (objective_id);
        """
    )
    c.commit()
    c.close()


def cmd_setup(args):
    existing = load_config()
    if existing and not args.force:
        print(
            json.dumps(
                {
                    "ok": False,
                    "error": "Already configured. Re-run with --force to change the exam date.",
                    "current": existing,
                },
                indent=2,
            )
        )
        sys.exit(1)

    exam = parse_exam_date(args.exam_date)
    delta = (exam - date.today()).days
    if delta < 0:
        print(
            json.dumps(
                {"ok": False, "error": f"{exam.isoformat()} is in the past. Confirm the date with the user."},
                indent=2,
            )
        )
        sys.exit(1)

    cfg = {
        "exam": "AB-100",
        "exam_name": "Agentic AI Business Solutions Architect",
        "exam_date": exam.isoformat(),
        "candidate_name": args.name or (existing or {}).get("candidate_name"),
        "daily_target": args.daily_target,
        "pass_score": 700,
        "configured_at": datetime.now().isoformat(timespec="seconds"),
    }
    save_config(cfg)
    ensure_schema()

    print(
        json.dumps(
            {
                "ok": True,
                "config": cfg,
                "days_until_exam": delta,
                "suggested_sessions_remaining": max(delta, 0),
                "message": "Setup complete. Run `python drill.py start` to begin the first drill.",
            },
            indent=2,
        )
    )


def cmd_status(_args):
    cfg = load_config()
    if cfg is None:
        print(json.dumps({"configured": False, "needs_onboarding": True, "action": ONBOARDING_MESSAGE}, indent=2))
        return
    sessions = 0
    answered = 0
    if os.path.exists(DB):
        ensure_schema()
        c = conn()
        sessions = c.execute("SELECT COUNT(*) n FROM sessions WHERE completed_at IS NOT NULL").fetchone()["n"]
        answered = c.execute("SELECT COUNT(*) n FROM answers").fetchone()["n"]
        c.close()
    print(
        json.dumps(
            {
                "configured": True,
                "config": cfg,
                "days_until_exam": days_until(cfg),
                "sessions_completed": sessions,
                "questions_answered": answered,
                "db_exists": os.path.exists(DB),
            },
            indent=2,
        )
    )


def cmd_init(_args):
    ensure_schema()
    print(json.dumps({"ok": True, "db": DB, "message": "Schema ready."}, indent=2))


def objective_stats(c):
    rows = c.execute(
        """
        SELECT objective_id,
               COUNT(*) AS attempts,
               SUM(correct) AS hits,
               MAX(session_id) AS last_session
        FROM answers GROUP BY objective_id
        """
    ).fetchall()
    return {r["objective_id"]: dict(r) for r in rows}


def allocate(count, weights, day_index):
    """Largest-remainder allocation with day-based tiebreak rotation."""
    keys = list(weights.keys())
    raw = {k: count * weights[k] / 100.0 for k in keys}
    alloc = {k: int(raw[k]) for k in keys}
    left = count - sum(alloc.values())
    order = sorted(keys, key=lambda k: (-(raw[k] - int(raw[k])), keys.index(k)))
    rot = day_index % max(len(order), 1)
    order = order[rot:] + order[:rot]
    for i in range(left):
        alloc[order[i % len(order)]] += 1
    return alloc


def cmd_start(args):
    cfg = require_config()
    bp = load_blueprint()
    ensure_schema()
    c = conn()
    today = date.today().isoformat()
    prior = c.execute("SELECT COUNT(*) n FROM sessions").fetchone()["n"]
    cur = c.execute(
        "INSERT INTO sessions (drill_date, started_at, planned_count) VALUES (?,?,?)",
        (today, datetime.now().isoformat(timespec="seconds"), args.count),
    )
    session_id = cur.lastrowid
    c.commit()

    stats = objective_stats(c)
    max_sess = c.execute("SELECT COALESCE(MAX(session_id),0) m FROM answers").fetchone()["m"]

    weights = {k: v["weight"] for k, v in bp["domains"].items()}
    alloc = allocate(args.count, weights, prior)

    rng = random.Random(session_id * 7919 + 13)
    picks = []
    for dom, n in alloc.items():
        pool = [o for o in bp["objectives"] if o["domain"] == dom]
        scored = []
        for o in pool:
            s = stats.get(o["id"])
            w = 1.0
            if s and s["attempts"]:
                acc = s["hits"] / s["attempts"]
                w += 2.6 * (1.0 - acc)
                gap = max_sess - s["last_session"]
                if gap <= 0:
                    w *= 0.30
                elif gap == 1:
                    w *= 0.60
            else:
                w += 1.3
            scored.append((o, max(w, 0.05)))
        chosen = []
        avail = scored[:]
        for _ in range(min(n, len(avail))):
            total = sum(w for _, w in avail)
            r = rng.uniform(0, total)
            acc = 0.0
            for i, (o, w) in enumerate(avail):
                acc += w
                if acc >= r:
                    chosen.append(o)
                    avail.pop(i)
                    break
        for o in chosen:
            s = stats.get(o["id"], {})
            att = s.get("attempts", 0) or 0
            hit = s.get("hits", 0) or 0
            picks.append(
                {
                    "objective_id": o["id"],
                    "domain": o["domain"],
                    "group": o["group"],
                    "objective": o["text"],
                    "prior_attempts": att,
                    "prior_accuracy": (round(hit / att, 2) if att else None),
                }
            )

    rng.shuffle(picks)
    for i, p in enumerate(picks, start=1):
        p["qnum"] = i

    recent = [
        r["stem"]
        for r in c.execute(
            "SELECT stem FROM answers WHERE stem IS NOT NULL ORDER BY id DESC LIMIT 60"
        ).fetchall()
        if r["stem"]
    ]
    c.close()

    print(
        json.dumps(
            {
                "session_id": session_id,
                "drill_date": today,
                "day_number": prior + 1,
                "days_until_exam": days_until(cfg),
                "exam_date": cfg["exam_date"],
                "candidate_name": cfg.get("candidate_name"),
                "planned_count": args.count,
                "domain_allocation": alloc,
                "questions": picks,
                "avoid_repeating_these_recent_stems": recent,
            },
            indent=2,
        )
    )


def cmd_record(args):
    bp = load_blueprint()
    dom = next((o["domain"] for o in bp["objectives"] if o["id"] == args.objective), None)
    if dom is None:
        print(json.dumps({"ok": False, "error": f"Unknown objective {args.objective}"}))
        sys.exit(1)
    ensure_schema()
    c = conn()
    c.execute(
        """INSERT OR REPLACE INTO answers
           (session_id, qnum, objective_id, domain, correct, stem, answered_at)
           VALUES (?,?,?,?,?,?,?)""",
        (
            args.session,
            args.qnum,
            args.objective,
            dom,
            1 if args.correct else 0,
            args.stem,
            datetime.now().isoformat(timespec="seconds"),
        ),
    )
    c.commit()
    row = c.execute(
        "SELECT COUNT(*) n, COALESCE(SUM(correct),0) k FROM answers WHERE session_id=?",
        (args.session,),
    ).fetchone()
    c.close()
    print(
        json.dumps(
            {"ok": True, "session_answered": row["n"], "session_correct": row["k"]}, indent=2
        )
    )


def cmd_finish(args):
    ensure_schema()
    c = conn()
    c.execute(
        "UPDATE sessions SET completed_at=? WHERE id=?",
        (datetime.now().isoformat(timespec="seconds"), args.session),
    )
    c.commit()
    c.close()
    print(json.dumps({"ok": True, "session_id": args.session}, indent=2))


def band(pct):
    if pct >= 85:
        return "STRONG - comfortably ready"
    if pct >= 75:
        return "READY - on track to pass"
    if pct >= 65:
        return "BORDERLINE - keep drilling weak areas"
    return "NOT READY - significant gaps remain"


def cmd_report(_args):
    cfg = require_config()
    bp = load_blueprint()
    ensure_schema()
    c = conn()
    total = c.execute("SELECT COUNT(*) n, COALESCE(SUM(correct),0) k FROM answers").fetchone()
    if not total["n"]:
        print(
            json.dumps(
                {
                    "message": "No answers recorded yet.",
                    "exam_date": cfg["exam_date"],
                    "days_until_exam": days_until(cfg),
                },
                indent=2,
            )
        )
        c.close()
        return

    per_domain = {}
    weighted = 0.0
    wsum = 0.0
    for dom, meta in bp["domains"].items():
        r = c.execute(
            "SELECT COUNT(*) n, COALESCE(SUM(correct),0) k FROM answers WHERE domain=?", (dom,)
        ).fetchone()
        acc = (r["k"] / r["n"] * 100) if r["n"] else None
        per_domain[dom] = {
            "name": meta["name"],
            "exam_weight": meta["range"],
            "answered": r["n"],
            "correct": r["k"],
            "accuracy_pct": (round(acc, 1) if acc is not None else None),
        }
        if acc is not None:
            weighted += acc * meta["weight"]
            wsum += meta["weight"]
    weighted_acc = round(weighted / wsum, 1) if wsum else 0.0

    obj_text = {o["id"]: o["text"] for o in bp["objectives"]}
    weak = [
        {
            "objective_id": r["objective_id"],
            "objective": obj_text.get(r["objective_id"], ""),
            "attempts": r["n"],
            "accuracy_pct": round(r["k"] / r["n"] * 100, 1),
        }
        for r in c.execute(
            """SELECT objective_id, COUNT(*) n, COALESCE(SUM(correct),0) k
               FROM answers GROUP BY objective_id
               HAVING k*1.0/n < 0.75
               ORDER BY k*1.0/n ASC, n DESC LIMIT 12"""
        ).fetchall()
    ]

    covered = c.execute("SELECT COUNT(DISTINCT objective_id) n FROM answers").fetchone()["n"]
    sessions = c.execute("SELECT COUNT(*) n FROM sessions WHERE completed_at IS NOT NULL").fetchone()["n"]
    c.close()

    print(
        json.dumps(
            {
                "exam": "AB-100",
                "candidate_name": cfg.get("candidate_name"),
                "exam_date": cfg["exam_date"],
                "days_until_exam": days_until(cfg),
                "sessions_completed": sessions,
                "questions_answered": total["n"],
                "raw_accuracy_pct": round(total["k"] / total["n"] * 100, 1),
                "blueprint_weighted_accuracy_pct": weighted_acc,
                "readiness_band": band(weighted_acc),
                "objective_coverage": f"{covered}/{len(bp['objectives'])}",
                "by_domain": per_domain,
                "weakest_objectives": weak,
            },
            indent=2,
        )
    )


def cmd_history(args):
    ensure_schema()
    c = conn()
    rows = c.execute(
        """SELECT s.id, s.drill_date, s.completed_at,
                  COUNT(a.id) n, COALESCE(SUM(a.correct),0) k
           FROM sessions s LEFT JOIN answers a ON a.session_id = s.id
           GROUP BY s.id ORDER BY s.id DESC LIMIT ?""",
        (args.limit,),
    ).fetchall()
    c.close()
    print(
        json.dumps(
            [
                {
                    "session_id": r["id"],
                    "date": r["drill_date"],
                    "answered": r["n"],
                    "correct": r["k"],
                    "accuracy_pct": (round(r["k"] / r["n"] * 100, 1) if r["n"] else None),
                    "completed": bool(r["completed_at"]),
                }
                for r in rows
            ],
            indent=2,
        )
    )


def cmd_reset(args):
    if not args.confirm:
        print(json.dumps({"ok": False, "error": "Pass --confirm to wipe progress.db and config.json."}, indent=2))
        sys.exit(1)
    removed = []
    for path in (DB, CONFIG):
        if os.path.exists(path):
            os.remove(path)
            removed.append(os.path.basename(path))
    print(json.dumps({"ok": True, "removed": removed, "message": "Run `setup` again to reconfigure."}, indent=2))


def main():
    p = argparse.ArgumentParser(description="AB-100 drill engine")
    sub = p.add_subparsers(dest="cmd", required=True)

    su = sub.add_parser("setup", help="First-run onboarding: set the target exam date")
    su.add_argument("--exam-date", required=True, help="ISO date, e.g. 2026-11-04")
    su.add_argument("--name", default=None)
    su.add_argument("--daily-target", type=int, default=20)
    su.add_argument("--force", action="store_true", help="Overwrite an existing exam date")
    su.set_defaults(func=cmd_setup)

    sub.add_parser("status").set_defaults(func=cmd_status)
    sub.add_parser("init").set_defaults(func=cmd_init)

    s = sub.add_parser("start")
    s.add_argument("--count", type=int, default=20)
    s.set_defaults(func=cmd_start)

    r = sub.add_parser("record")
    r.add_argument("--session", type=int, required=True)
    r.add_argument("--qnum", type=int, required=True)
    r.add_argument("--objective", required=True)
    r.add_argument("--correct", type=int, required=True)
    r.add_argument("--stem", default=None)
    r.set_defaults(func=cmd_record)

    f = sub.add_parser("finish")
    f.add_argument("--session", type=int, required=True)
    f.set_defaults(func=cmd_finish)

    sub.add_parser("report").set_defaults(func=cmd_report)

    h = sub.add_parser("history")
    h.add_argument("--limit", type=int, default=10)
    h.set_defaults(func=cmd_history)

    rs = sub.add_parser("reset")
    rs.add_argument("--confirm", action="store_true")
    rs.set_defaults(func=cmd_reset)

    args = p.parse_args()
    args.func(args)


if __name__ == "__main__":
    main()
