"""
simulate_days.py — large day/stage simulation. Takes a diverse set of real questions (KB questions
+ real logged coach-test messages) and runs EACH across the program+day stages, then auto-grades:
  - outcome distribution (answered / escalated / declined) overall and per stage
  - program leaks (a Relaxed answer asserting Original's "Day 6 quit", or vice versa)
  - weak-confidence answers (answered but low retrieval similarity)
  - early-stage concept mentions (mindful smoking / deceptions surfaced to very early users)
  - faithfulness on a random sample (LLM judge vs the retrieved approved sources)

Run:  python3 eval/simulate_days.py   (writes eval/sim_days_report.md)
"""
import os, sys, json, re, random
from collections import defaultdict, Counter
from concurrent.futures import ThreadPoolExecutor

ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, ROOT)
random.seed(42)

import app  # noqa: E402
from app import ChatRequest, retrieve, client  # noqa: E402
from google.genai import types  # noqa: E402

# --- 1. build a diverse base question set: KB questions + real logged user messages
KB = [e for e in json.load(open(os.path.join(ROOT, "data", "kb.json")))
      if not e.get("coach_only") and not e.get("has_placeholder")]
kb_qs = [e["question"] for e in KB]

logged = []
try:
    import pymysql
    host = os.environ.get("DB_HOST") or os.environ.get("DB_HOST_DISABLED")
    if host:
        cn = pymysql.connect(host=host, port=int(os.environ.get("DB_PORT", 3306)),
                             user=os.environ["DB_USERNAME"], password=os.environ["DB_PASSWORD"],
                             database=os.environ["DB_DATABASE"], cursorclass=pymysql.cursors.DictCursor)
        cur = cn.cursor()
        cur.execute("SELECT DISTINCT vMessage m FROM tbl_SupportChatLog "
                    "WHERE dCreated >= '2026-07-07' AND CHAR_LENGTH(vMessage) BETWEEN 12 AND 220")
        logged = [r["m"] for r in cur.fetchall() if r["m"]]
        cn.close()
except Exception as e:
    print("(no DB questions:", str(e)[:60], ")")

pool = list(dict.fromkeys([q.strip() for q in (logged + kb_qs) if q and q.strip()]))
random.shuffle(pool)
BASE = pool[:130]  # 130 base x 8 stages ~= 1040 runs
print(f"base questions: {len(BASE)} (logged={len(logged)}, kb={len(kb_qs)})")

# --- 2. stage combos (representative day inside each range)
STAGES = [("P3", 2, "P3 d1-2 intro"), ("P3", 4, "P3 d3-5 content"), ("P3", 6, "P3 d6 quit"),
          ("P3", 10, "P3 post-quit"), ("P9", 6, "P9 d1-11 early"), ("P9", 20, "P9 d12-30 mid"),
          ("P9", 38, "P9 d31-42 late"), ("P9", 45, "P9 post-quit")]

RUNS = [(q, prog, day, label) for q in BASE for (prog, day, label) in STAGES]
print(f"total runs: {len(RUNS)}")

# leak patterns
ORIG_LEAK = re.compile(r"quit[^.]{0,20}day 6|day 6[^.]{0,20}quit|6[- ]day program|only 6 days|"
                       r"complete[^.]{0,15}\b6 days\b|in just 6 days", re.I)
RELX_LEAK = re.compile(r"\b42[- ]?day|\bday 42\b|quit[^.]{0,20}day 42", re.I)
CONCEPT = re.compile(r"mindful smoking|smoke mindfully|\bdeception", re.I)


def one(run):
    q, prog, day, label = run
    try:
        o = app.chat(ChatRequest(message=q, program=prog, program_day=day))
        _, best = retrieve(q, program=prog)
        a = o.get("answer") or ""
        esc = o.get("escalate", False)
        reason = o.get("escalate_reason")
        outcome = ("declined" if reason == "out_of_scope" else
                   ("escalated" if esc else "answered"))
        leak = ""
        if not esc:
            if prog == "P9" and ORIG_LEAK.search(a):
                leak = "orig-in-relaxed"
            elif prog == "P3" and RELX_LEAK.search(a):
                leak = "relaxed-in-orig"
        early = (prog == "P3" and day <= 2) or (prog == "P9" and day <= 11)
        concept = bool(early and outcome == "answered" and CONCEPT.search(a))
        return {"q": q, "prog": prog, "day": day, "label": label, "outcome": outcome,
                "reason": reason, "best": round(best, 3), "leak": leak, "early_concept": concept,
                "answer": a}
    except Exception as e:
        return {"q": q, "prog": prog, "day": day, "label": label, "outcome": "error",
                "reason": "err", "best": 0.0, "leak": "", "early_concept": False,
                "answer": "ERR:" + str(e)[:60]}


def main():
    results = []
    with ThreadPoolExecutor(max_workers=10) as ex:
        for i, r in enumerate(ex.map(one, RUNS)):
            results.append(r)
            if (i + 1) % 100 == 0:
                print(f"  {i+1}/{len(RUNS)}")

    # --- faithfulness on a random sample of answered runs
    RULES = ("You grade whether a support answer is SUPPORTED by the approved sources. faithful=false "
             "if it adds any factual claim, step, or advice NOT in the sources (ignore warmth/empathy). "
             'Respond JSON {"faithful":bool,"why":str}')
    SCH = {"type": "object", "properties": {"faithful": {"type": "boolean"}, "why": {"type": "string"}},
           "required": ["faithful"]}
    answered = [r for r in results if r["outcome"] == "answered"]
    sample = random.sample(answered, min(60, len(answered)))

    def judge(r):
        hits, _ = retrieve(r["q"], program=r["prog"])
        src = "\n\n".join(f"{e['question']}: {e['answer']}" for e, s in hits if s > app.RETRIEVAL_FLOOR)
        if not src:
            return None
        try:
            resp = client.models.generate_content(
                model="gemini-2.5-flash",
                contents=f"QUESTION:\n{r['q']}\n\nAPPROVED SOURCES:\n{src}\n\nANSWER:\n{r['answer']}",
                config=types.GenerateContentConfig(system_instruction=RULES, response_mime_type="application/json",
                                                   response_schema=SCH, temperature=0.0))
            j = json.loads(resp.text)
            return (r, j)
        except Exception:
            return None
    faith = []
    with ThreadPoolExecutor(max_workers=8) as ex:
        for x in ex.map(judge, sample):
            if x:
                faith.append(x)
    unfaith = [(r, j) for r, j in faith if not j["faithful"]]

    # --- report
    L = ["# Day/stage simulation report\n"]
    L.append(f"Runs: {len(results)}  ({len(BASE)} base questions x {len(STAGES)} stages)\n")
    oc = Counter(r["outcome"] for r in results)
    L.append("## Outcomes (overall)")
    for k in ("answered", "escalated", "declined", "error"):
        L.append(f"- {k}: {oc.get(k,0)}  ({100*oc.get(k,0)//len(results)}%)")
    leaks = [r for r in results if r["leak"]]
    weak = [r for r in results if r["outcome"] == "answered" and r["best"] < 0.60]
    ec = [r for r in results if r["early_concept"]]
    L.append(f"\n**Program leaks: {len(leaks)}**   **Weak-confidence answers (sim<0.60): {len(weak)}**"
             f"   **Early-stage concept mentions: {len(ec)}**")
    if faith:
        L.append(f"**Faithfulness sample: {len(faith)-len(unfaith)}/{len(faith)} grounded "
                 f"({100*(len(faith)-len(unfaith))//len(faith)}%)**")

    L.append("\n## Outcomes by stage")
    by = defaultdict(Counter)
    for r in results:
        by[r["label"]][r["outcome"]] += 1
    L.append("| stage | answered | escalated | declined | err |")
    L.append("|---|---|---|---|---|")
    for _, _, label in STAGES:
        c = by[label]
        L.append(f"| {label} | {c['answered']} | {c['escalated']} | {c['declined']} | {c['error']} |")

    if leaks:
        L.append("\n## PROGRAM LEAKS (review)")
        for r in leaks[:40]:
            L.append(f"- [{r['label']}] ({r['leak']}) {r['q'][:60]}")
    if weak:
        L.append("\n## WEAK-CONFIDENCE ANSWERS (sim<0.60, review)")
        for r in sorted(weak, key=lambda x: x["best"])[:30]:
            L.append(f"- sim={r['best']} [{r['label']}] {r['q'][:65]}")
    if ec:
        L.append("\n## EARLY-STAGE CONCEPT MENTIONS (mindful/deception to early users, review)")
        for r in ec[:30]:
            L.append(f"- [{r['label']}] {r['q'][:70]}")
    if unfaith:
        L.append("\n## OFF-SHEET (faithfulness sample)")
        for r, j in unfaith:
            L.append(f"- [{r['label']}] {r['q'][:55]}\n    -> {j.get('why','')[:150]}")

    open(os.path.join(ROOT, "eval", "sim_days_report.md"), "w").write("\n".join(L))
    json.dump(results, open(os.path.join(ROOT, "eval", "sim_days_results.json"), "w"), ensure_ascii=False)
    print("\n".join(L[:14]))
    print("\nFull report -> eval/sim_days_report.md")


if __name__ == "__main__":
    main()
