"""Phase 3 exclusion check: is anything about this study's results already public?

For each candidate id in the screened pool:
  1. Registry, live: CT.gov hasResults / RESULT references / status; OSF has_papers / article_doi / withdrawn.
  2. Europe PMC (journals, PubMed, preprint servers incl. medRxiv/bioRxiv/PsyArXiv, many abstracts):
     search for the registry id anywhere in the record, then for the exact title.
  3. OpenAlex: search for the registry id, then for the exact title.
Any hit is written out with its URL for a human or agent to read; a hit is not automatically an
exclusion (a protocol paper or a baseline paper does not reveal results). The decision rule:
  - results posted on the registry                                   -> EXCLUDE (results posted)
  - a hit whose title/type indicates results, preprint or abstract   -> REVIEW
  - a hit that is a protocol / design / rationale paper             -> keep, noted
  - no hits                                                          -> keep
Writes the output path given as argv[2]; resumable (see __main__).
"""
import json, os, re, sys, time, urllib.parse, urllib.request, datetime

TODAY = datetime.date.today().isoformat()
# OpenAlex without an API key has a small daily budget per IP; on 2026-09-30 it was exhausted
# (HTTP 429, "Insufficient budget"). SKIP_OPENALEX=1 runs registry + Europe PMC only and records
# the skip on every row in "sources_skipped", so the gap is visible, not silent.
SKIP_OPENALEX = os.environ.get("SKIP_OPENALEX") == "1"
UA = {"User-Agent": "pm-forecast-exclusion-check"}


def get(url):
    for i in range(3):
        try:
            return json.load(urllib.request.urlopen(urllib.request.Request(url, headers=UA), timeout=60))
        except urllib.error.HTTPError as e:
            if e.code == 404:
                return None
            err = e
            if e.code == 429:  # rate limit: back off hard rather than record a failed check
                time.sleep(30 * (i + 1)); continue
        except Exception as e:
            err = e
        time.sleep(2 * (i + 1))
    raise err


def europepmc(q):
    d = get("https://www.ebi.ac.uk/europepmc/webservices/rest/search?" +
            urllib.parse.urlencode({"query": q, "format": "json", "pageSize": 25, "resultType": "lite"}))
    out = []
    for r in (d or {}).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')}"
        out.append({"src": "europepmc", "title": r.get("title", ""), "url": url, "date": r.get("firstPublicationDate", ""),
                    "type": r.get("pubType", ""), "source": r.get("source", "")})
    return out


def openalex(q):
    d = get("https://api.openalex.org/works?" + urllib.parse.urlencode({"search": q, "per-page": 25}))
    return [{"src": "openalex", "title": w.get("title") or "", "url": w.get("doi") or w.get("id"),
             "date": w.get("publication_date", ""), "type": w.get("type", ""), "source": ""}
            for w in (d or {}).get("results", [])]


PROTOCOL_WORDS = re.compile(r"protocol|design and rationale|rationale and design|study design|statistical analysis plan|registered report: stage 1", re.I)


def classify(hits):
    if not hits:
        return "no hits"
    if all(PROTOCOL_WORDS.search(h["title"]) for h in hits):
        return "protocol paper only"
    return "hits to review"


def check(c):
    sid = c["id"]
    res = {"id": sid, "checked_on": TODAY, "registry": {}, "hits": [], "error": ""}
    try:
        if sid.startswith("NCT"):
            d = get(f"https://clinicaltrials.gov/api/v2/studies/{sid}")
            p = d["protocolSection"]
            refs = [r for r in p.get("referencesModule", {}).get("references", []) if r.get("type") == "RESULT"]
            res["registry"] = {"status": p["statusModule"]["overallStatus"], "hasResults": d.get("hasResults", False),
                               "result_refs": [r.get("citation", "")[:200] for r in refs],
                               "primary_completion": p["statusModule"].get("primaryCompletionDateStruct", {}).get("date")}
            key = sid
        elif not sid.startswith("osf-"):
            # Registered Report / SSPP record: no registry API; search by title and URL only
            res["registry"] = {"status": "IPA" if c.get("source") != "sspp" else "SSPP-open", "hasResults": False, "result_refs": []}
            key = c.get("url", sid)
        else:
            oid = sid[4:]
            a = get(f"https://api.osf.io/v2/registrations/{oid}/")["data"]["attributes"]
            res["registry"] = {"status": "WITHDRAWN" if a.get("withdrawn") else "REGISTERED",
                               "hasResults": bool(a.get("has_papers") or a.get("article_doi")),
                               "result_refs": [a["article_doi"]] if a.get("article_doi") else []}
            key = f"osf.io/{oid}"
        title = c.get("title", "")[:250]
        hits, errs = [], []
        queries = [(europepmc, f'"{key}"'), (openalex, key)] if sid.startswith(("NCT", "osf-")) else []
        if title:
            t = re.sub(r'[^\w\s-]', " ", title)
            queries += [(europepmc, f'TITLE:"{t}"'), (openalex, t)]
        if SKIP_OPENALEX:
            queries = [q for q in queries if q[0] is not openalex]
            res["sources_skipped"] = "openalex (daily budget exhausted)"
        for fn, q in queries:
            try:
                hits += fn(q)
            except Exception as e:
                errs.append(f"{fn.__name__}: {e!r}"[:150])
            time.sleep(1.0)
        if len(errs) == len(queries):
            raise RuntimeError("; ".join(errs))
        res["partial_errors"] = errs
        # keep only openalex title-search hits that actually share most of the title words
        tw = set(re.findall(r"[a-z]{4,}", title.lower()))
        kept = []
        for h in hits:
            u = (h["url"] or "").lower()
            if "osf.io" in u or "10.17605/osf" in u or "clinicaltrials.gov" in u:
                continue  # the registration's own record, not a publication
            hw = set(re.findall(r"[a-z]{4,}", h["title"].lower()))
            if key.lower() in json.dumps(h).lower() or (tw and len(tw & hw) / max(1, len(tw)) >= 0.6):
                kept.append(h)
        seen, res["hits"] = set(), []
        for h in kept:
            if h["url"] not in seen:
                seen.add(h["url"]); res["hits"].append(h)
    except Exception as e:
        res["error"] = repr(e)[:300]
    r = res["registry"]
    if r.get("hasResults") or r.get("result_refs"):
        res["decision"] = "EXCLUDE: results or result reference on registry"
    elif r.get("status") in ("TERMINATED", "WITHDRAWN", "SUSPENDED"):
        res["decision"] = "EXCLUDE: registry status " + r["status"]
    elif res["error"] or res.get("partial_errors"):
        res["decision"] = "REVIEW: check failed"
    else:
        cl = classify(res["hits"])
        res["decision"] = {"no hits": "KEEP: no publication found", "protocol paper only": "KEEP: protocol paper only",
                           "hits to review": "REVIEW: publications found"}[cl]
    return res


if __name__ == "__main__":
    # Resumable: results already in the output file are kept; rows whose check failed are redone.
    # The file is rewritten after every study so an interrupted run loses nothing.
    import os
    pool = json.load(open(sys.argv[1]))
    done = {o["id"]: o for o in json.load(open(sys.argv[2]))} if os.path.exists(sys.argv[2]) else {}
    for i, c in enumerate(pool, 1):
        if c["id"] not in done or done[c["id"]]["decision"] == "REVIEW: check failed":
            done[c["id"]] = check(c)
            json.dump([done[x["id"]] for x in pool if x["id"] in done], open(sys.argv[2], "w"), indent=1)
            time.sleep(0.3)
        if i % 25 == 0:
            print(i, "checked", flush=True)
    out = [done[x["id"]] for x in pool]
    import collections
    print(collections.Counter(o["decision"].split(":")[0] for o in out))
