#!/usr/bin/env python3
"""Phase 6: resolve forecast studies and score FRAMEWORK vs BASELINE.

Two subcommands:

  python3 scripts/resolve.py check
      For every study in data/studies.csv, look for results:
        - ClinicalTrials.gov: overall status, hasResults, and references marked RESULT
        - Europe PMC: full-text search for the registry id (catches papers and preprints)
        - OSF: registration has_papers / article_doi
      Writes data/resolution_log.csv (one row per study per run) and appends any study with
      evidence of results to data/resolution_queue.csv. It never writes an outcome by itself:
      deciding whether a primary hypothesis was supported means reading the paper, so a
      queued study is adjudicated by hand (or by a Claude session) into data/outcomes.csv.
      A study whose registry status becomes TERMINATED / WITHDRAWN / SUSPENDED is written to
      outcomes.csv as status=void and dropped from scoring (reported, not hidden).

  python3 scripts/resolve.py score
      Reads data/outcomes.csv and forecasts/*/*.json and writes reports/scores.md and
      data/claim_track_record.csv.

outcomes.csv columns:
  study_id, status (resolved | void), primary_supported (1 | 0 | blank),
  observed_effect (number, in the metric named), effect_metric,
  moderator_reported (yes | no), moderator_observed (free text: which subgroup responded more),
  sspp_crowd_p (0-1, blank if none), source_url, adjudicated_by, adjudicated_on, notes
"""
import csv, glob, json, math, os, random, sys, time, urllib.parse, urllib.request, datetime

ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
os.chdir(ROOT)
TODAY = datetime.date.today().isoformat()
OUT_COLS = ["study_id", "status", "primary_supported", "observed_effect", "effect_metric", "moderator_reported",
            "moderator_observed", "sspp_crowd_p", "source_url", "adjudicated_by", "adjudicated_on", "notes"]


def get_json(url):
    req = urllib.request.Request(url, headers={"User-Agent": "pm-forecast-resolve"})
    for i in range(3):
        try:
            return json.load(urllib.request.urlopen(req, timeout=60))
        except Exception as e:
            err = e
            time.sleep(3 * (i + 1))
    raise err


def read_csv(path):
    return list(csv.DictReader(open(path, newline=""))) if os.path.exists(path) else []


def append_csv(path, cols, rows):
    new = not os.path.exists(path)
    with open(path, "a", newline="") as f:
        w = csv.DictWriter(f, fieldnames=cols)
        if new:
            w.writeheader()
        w.writerows(rows)


# ---------------------------------------------------------------- check

def check_ctgov(nct):
    d = get_json(f"https://clinicaltrials.gov/api/v2/studies/{nct}")
    p = d["protocolSection"]
    refs = [r for r in p.get("referencesModule", {}).get("references", []) if r.get("type") == "RESULT"]
    return {
        "registry_status": p["statusModule"].get("overallStatus"),
        "has_results": d.get("hasResults", False),
        "result_refs": ["https://pubmed.ncbi.nlm.nih.gov/%s/" % r["pmid"] for r in refs if r.get("pmid")],
    }


def check_osf(osf_id):
    d = get_json(f"https://api.osf.io/v2/registrations/{osf_id}/")["data"]["attributes"]
    return {"registry_status": "WITHDRAWN" if d.get("withdrawn") else "REGISTERED",
            "has_results": bool(d.get("has_papers") or d.get("article_doi")),
            "result_refs": ["https://doi.org/" + d["article_doi"]] if d.get("article_doi") else []}


def check_europepmc(registry_id):
    q = urllib.parse.quote(f'"{registry_id}"')
    d = get_json(f"https://www.ebi.ac.uk/europepmc/webservices/rest/search?query={q}&format=json&pageSize=25")
    hits = []
    for r in d.get("resultList", {}).get("result", []):
        url = ("https://doi.org/" + r["doi"]) if r.get("doi") else f"https://europepmc.org/article/{r.get('source')}/{r.get('id')}"
        hits.append({"url": url, "title": r.get("title", ""), "pubtype": r.get("pubType", ""), "date": r.get("firstPublicationDate", "")})
    return hits


def cmd_check():
    studies = read_csv("data/studies.csv")
    done = {r["study_id"] for r in read_csv("data/outcomes.csv")}
    log_rows, queue_rows, void_rows = [], [], []
    for s in studies:
        sid = s["study_id"]
        if sid in done:
            continue
        row = {"study_id": sid, "checked_on": TODAY, "registry_status": "", "has_results": "", "evidence": "", "error": ""}
        try:
            reg = check_osf(sid[4:]) if sid.startswith("osf-") else check_ctgov(sid) if sid.startswith("NCT") else None
            ids = [sid[4:] if sid.startswith("osf-") else sid]
            hits = []
            for i in ids:
                hits += check_europepmc(i)
            ev = (reg["result_refs"] if reg else []) + [h["url"] for h in hits]
            row.update(registry_status=reg["registry_status"] if reg else "n/a",
                       has_results=str(reg["has_results"]) if reg else "n/a", evidence=" ".join(ev))
            if reg and reg["registry_status"] in ("TERMINATED", "WITHDRAWN", "SUSPENDED"):
                void_rows.append({"study_id": sid, "status": "void", "source_url": s["registry_url"],
                                  "adjudicated_by": "resolve.py", "adjudicated_on": TODAY,
                                  "notes": "registry status " + reg["registry_status"]})
            elif (reg and reg["has_results"]) or ev:
                queue_rows.append({"study_id": sid, "found_on": TODAY, "registry_url": s["registry_url"],
                                   "evidence": " ".join(ev), "hits": json.dumps(hits)[:4000]})
        except Exception as e:
            row["error"] = repr(e)[:300]
        log_rows.append(row)
        time.sleep(0.4)
    append_csv("data/resolution_log.csv", list(log_rows[0].keys()) if log_rows else ["study_id"], log_rows)
    if queue_rows:
        append_csv("data/resolution_queue.csv", list(queue_rows[0].keys()), queue_rows)
    if void_rows:
        append_csv("data/outcomes.csv", OUT_COLS, [{c: r.get(c, "") for c in OUT_COLS} for r in void_rows])
    errs = sum(1 for r in log_rows if r["error"])
    print(f"checked {len(log_rows)} studies: {len(queue_rows)} queued for adjudication, "
          f"{len(void_rows)} void, {errs} errors (see data/resolution_log.csv)")


# ---------------------------------------------------------------- score

def load_forecasts():
    out = {}
    for path in glob.glob("forecasts/*/*.json"):
        sid = os.path.basename(os.path.dirname(path))
        agent = os.path.basename(path).rsplit("_", 1)[0]
        try:
            out.setdefault(sid, {}).setdefault(agent, []).append(json.load(open(path)))
        except Exception:
            pass
    return out


def interval(f):
    iv = f.get("effect_size_90pct_interval")
    if isinstance(iv, dict):
        lo, hi = iv.get("low", iv.get("lower")), iv.get("high", iv.get("upper"))
    elif isinstance(iv, (list, tuple)) and len(iv) == 2:
        lo, hi = iv
    else:
        return None
    try:
        return float(lo), float(hi)
    except (TypeError, ValueError):
        return None


def bootstrap_ci(diffs, n=10000, seed=7):
    rnd = random.Random(seed)
    means = sorted(sum(rnd.choice(diffs) for _ in diffs) / len(diffs) for _ in range(n))
    return means[int(0.025 * n)], means[int(0.975 * n)]


def signed_rank_p(diffs):
    try:
        from scipy.stats import wilcoxon
        nz = [d for d in diffs if d != 0]
        return wilcoxon(nz).pvalue if len(nz) >= 5 else float("nan")
    except Exception:
        return float("nan")


def cmd_score():
    outcomes = {r["study_id"]: r for r in read_csv("data/outcomes.csv") if r["status"] == "resolved"}
    voids = [r for r in read_csv("data/outcomes.csv") if r["status"] == "void"]
    fc = load_forecasts()
    rows, claim_rec = [], {}
    for sid, o in outcomes.items():
        if sid not in fc or o["primary_supported"] not in ("0", "1"):
            continue
        y = int(o["primary_supported"])
        rec = {"study_id": sid}
        for agent in ("baseline", "framework"):
            runs = fc[sid].get(agent, [])
            ps = [float(r["p_primary_hypothesis_supported"]) for r in runs if "p_primary_hypothesis_supported" in r]
            if not ps:
                continue
            p = sum(ps) / len(ps)
            rec[agent + "_p"] = p
            rec[agent + "_brier"] = (p - y) ** 2
            if o.get("observed_effect"):
                obs = float(o["observed_effect"])
                ivs = [interval(r) for r in runs]
                ivs = [i for i in ivs if i]
                if ivs:
                    rec[agent + "_coverage"] = sum(lo <= obs <= hi for lo, hi in ivs) / len(ivs)
                    rec[agent + "_abs_err"] = sum(abs((lo + hi) / 2 - obs) for lo, hi in ivs) / len(ivs)
            if o.get("moderator_reported") == "yes":
                # a moderator hit is judged by hand into the forecast's run file as "moderator_hit": true/false
                hits = [r["moderator_hit"] for r in runs if isinstance(r.get("moderator_hit"), bool)]
                if hits:
                    rec[agent + "_moderator_hit"] = sum(hits) / len(hits)
            if agent == "framework":
                for r in runs:
                    for cid in r.get("claim_ids_used", []):
                        c = claim_rec.setdefault(cid, {"claim_id": cid, "n_runs": 0, "studies": set(), "brier_sum": 0.0})
                        c["n_runs"] += 1
                        c["studies"].add(sid)
                        c["brier_sum"] += (float(r["p_primary_hypothesis_supported"]) - y) ** 2
        if o.get("sspp_crowd_p"):
            rec["crowd_brier"] = (float(o["sspp_crowd_p"]) - y) ** 2
        rows.append(rec)

    lines = [f"# Forecast scores ({TODAY})", "",
             f"Resolved studies: {len(rows)}. Void (terminated/withdrawn, excluded from scoring): {len(voids)}.", ""]
    paired = [r for r in rows if "baseline_brier" in r and "framework_brier" in r]
    if paired:
        d = [r["framework_brier"] - r["baseline_brier"] for r in paired]
        mb = sum(r["baseline_brier"] for r in paired) / len(paired)
        mf = sum(r["framework_brier"] for r in paired) / len(paired)
        lo, hi = bootstrap_ci(d) if len(d) > 1 else (float("nan"), float("nan"))
        lines += ["## Brier score (lower is better)", "",
                  "| | n | mean Brier |", "|---|---|---|",
                  f"| BASELINE | {len(paired)} | {mb:.4f} |", f"| FRAMEWORK | {len(paired)} | {mf:.4f} |", "",
                  f"FRAMEWORK minus BASELINE: {sum(d)/len(d):+.4f}, 95% bootstrap CI [{lo:+.4f}, {hi:+.4f}], "
                  f"Wilcoxon signed-rank p = {signed_rank_p(d):.3g}. A negative difference favours the framework.", ""]
        crowd = [r for r in paired if "crowd_brier" in r]
        if crowd:
            lines += [f"Third comparison, SSPP crowd (n={len(crowd)}): mean Brier "
                      f"{sum(r['crowd_brier'] for r in crowd)/len(crowd):.4f} vs FRAMEWORK "
                      f"{sum(r['framework_brier'] for r in crowd)/len(crowd):.4f} and BASELINE "
                      f"{sum(r['baseline_brier'] for r in crowd)/len(crowd):.4f} on the same studies.", ""]
    for metric, label in (("coverage", "90% interval coverage (target 0.90)"), ("abs_err", "Absolute error of interval midpoint"),
                          ("moderator_hit", "Moderator hit rate")):
        pr = [r for r in rows if f"baseline_{metric}" in r and f"framework_{metric}" in r]
        if pr:
            b = sum(r[f"baseline_{metric}"] for r in pr) / len(pr)
            f_ = sum(r[f"framework_{metric}"] for r in pr) / len(pr)
            lines += [f"## {label}", "", f"n = {len(pr)} studies. BASELINE {b:.3f}, FRAMEWORK {f_:.3f}.", ""]
    lines += ["## Per study", "", "| study | outcome | BASELINE p | FRAMEWORK p | BASELINE Brier | FRAMEWORK Brier |",
              "|---|---|---|---|---|---|"]
    for r in sorted(rows, key=lambda r: r["study_id"]):
        o = outcomes[r["study_id"]]
        lines.append(f"| {r['study_id']} | {'supported' if o['primary_supported']=='1' else 'not supported'} | "
                     f"{r.get('baseline_p', float('nan')):.2f} | {r.get('framework_p', float('nan')):.2f} | "
                     f"{r.get('baseline_brier', float('nan')):.3f} | {r.get('framework_brier', float('nan')):.3f} |")
    lines += ["", "Misses are listed with the same weight as hits: every resolved study is in the table above."]
    os.makedirs("reports", exist_ok=True)
    open("reports/scores.md", "w").write("\n".join(lines) + "\n")

    with open("data/claim_track_record.csv", "w", newline="") as f:
        w = csv.writer(f)
        w.writerow(["claim_id", "n_studies", "n_runs", "mean_brier_of_runs_citing_it"])
        for c in sorted(claim_rec.values(), key=lambda c: c["claim_id"]):
            w.writerow([c["claim_id"], len(c["studies"]), c["n_runs"], f"{c['brier_sum']/c['n_runs']:.4f}"])
    print(f"scored {len(rows)} resolved studies; wrote reports/scores.md and data/claim_track_record.csv")


if __name__ == "__main__":
    cmd = sys.argv[1] if len(sys.argv) > 1 else ""
    if cmd == "check":
        cmd_check()
    elif cmd == "score":
        cmd_score()
    else:
        print(__doc__)
