"""Family-level colexification rates in CLICS4 v1.0 (see plan/analysis-plan.md).

Two concepts are colexified in a language variety when the variety lists an identical form
(after Unicode NFC normalisation, case-folding and trimming) for both. For a concept pair,
the family-level rate is the share of language families, among families in which at least
one variety has forms for both concepts, in which at least one variety colexifies them.

Usage: python code/colex.py <forms.csv> <languages.csv> <out.json> [--check PAIR ...]
"""
import csv
import json
import sys
import unicodedata
from collections import defaultdict

import numpy as np

TARGETS = ["SAD", "GRIEF", "DIFFICULT", "TIRED"]
PHYSICAL = ["HEAVY", "LIGHT (WEIGHT)", "LOW", "DOWN", "COLD", "HOT", "DARK", "BITTER", "SWEET",
            "HARD", "SOFT", "THICK", "DEEP", "WEAK", "SLOW"]
GRAVITY = ["HEAVY", "LOW", "DOWN"]
AFFECT = ["SAD", "GRIEF"]
MIN_FAMILIES = 30
SEED = 20261007
N_PERM = 10000


def norm(s):
    return unicodedata.normalize("NFC", s).casefold().strip()


def load(forms_path, langs_path):
    fam = {}
    with open(langs_path, newline="", encoding="utf-8") as fh:
        for r in csv.DictReader(fh):
            fam[r["ID"]] = r["Family"] or ("isolate:" + r["ID"])
    gloss_of = {}
    with open(forms_path.replace("forms.csv", "concepts.csv"), newline="", encoding="utf-8") as fh:
        for r in csv.DictReader(fh):
            gloss_of[r["ID"]] = r["Concepticon_Gloss"]
    # per variety: concept -> set of forms
    lex = defaultdict(lambda: defaultdict(set))
    with open(forms_path, newline="", encoding="utf-8") as fh:
        for r in csv.DictReader(fh):
            g = gloss_of.get(r["Parameter_ID"])
            f = norm(r["Form"] or "")
            if g and f:
                lex[r["Language_ID"]][g].add(f)
    return lex, fam


def rates_for(anchor, lex, fam):
    """Family-level colexification rate of anchor with every other concept."""
    both = defaultdict(set)   # concept -> families where some variety has anchor and concept
    colex = defaultdict(set)  # concept -> families where some variety colexifies them
    for var, concepts in lex.items():
        if anchor not in concepts:
            continue
        af = concepts[anchor]
        f = fam.get(var, "unknown:" + var)
        for c, forms in concepts.items():
            if c == anchor:
                continue
            both[c].add(f)
            if af & forms:
                colex[c].add(f)
    return {c: (len(colex[c]), len(both[c])) for c in both}


def percentile_of(value, dist):
    dist = np.asarray(dist)
    return float((dist < value).mean() + 0.5 * (dist == value).mean())


def main():
    forms_path, langs_path, out = sys.argv[1], sys.argv[2], sys.argv[3]
    lex, fam = load(forms_path, langs_path)
    res = {"varieties": len(lex), "families": len(set(fam.values())), "min_families": MIN_FAMILIES}
    if "--check" in sys.argv:
        # pre-registration check on non-target pairs only
        pairs = sys.argv[sys.argv.index("--check") + 1:]
        for p in pairs:
            a, b = p.split("|")
            k, n = rates_for(a, lex, fam).get(b, (0, 0))
            print(p, k, n)
        return
    # H1: HEAVY with each target, against HEAVY's rates with all other eligible concepts
    hv = rates_for("HEAVY", lex, fam)
    elig = {c: v for c, v in hv.items() if v[1] >= MIN_FAMILIES}
    dist = [k / n for c, (k, n) in elig.items() if c not in TARGETS]
    h1 = {}
    for t in TARGETS:
        k, n = hv.get(t, (0, 0))
        rate = k / n if n else float("nan")
        h1[t] = {"families_colexifying": k, "families_with_both": n, "rate": round(rate, 4),
                 "percentile_among_heavy_partners": round(percentile_of(rate, dist), 4),
                 "top_5_percent": bool(percentile_of(rate, dist) >= 0.95)}
    res["H1_heavy"] = {"eligible_partners": len(dist), "median_partner_rate": round(float(np.median(dist)), 4),
                       "p95_partner_rate": round(float(np.percentile(dist, 95)), 4), "targets": h1}
    # H2: gravity-related versus other physical-property words, colexification with SAD or GRIEF
    tab = {}
    for p in PHYSICAL:
        r = rates_for(p, lex, fam)
        fams_both, fams_colex = set(), set()
        for var, concepts in lex.items():
            if p in concepts and any(a in concepts for a in AFFECT):
                f = fam.get(var, "unknown:" + var)
                fams_both.add(f)
                if any(concepts[p] & concepts[a] for a in AFFECT if a in concepts):
                    fams_colex.add(f)
        tab[p] = {"families_colexifying": len(fams_colex), "families_with_both": len(fams_both),
                  "rate": round(len(fams_colex) / len(fams_both), 4) if fams_both else float("nan")}
    res["H2_physical_with_sad_or_grief"] = tab
    g = np.array([tab[p]["rate"] for p in GRAVITY])
    o = np.array([tab[p]["rate"] for p in PHYSICAL if p not in GRAVITY])
    obs = float(g.mean() - o.mean())
    rng = np.random.default_rng(SEED)
    allr = np.array([tab[p]["rate"] for p in PHYSICAL])
    k = len(GRAVITY)
    perm = []
    for _ in range(N_PERM):
        idx = rng.permutation(len(allr))
        perm.append(allr[idx[:k]].mean() - allr[idx[k:]].mean())
    perm = np.array(perm)
    res["H2_test"] = {"gravity_mean_rate": round(float(g.mean()), 4), "other_mean_rate": round(float(o.mean()), 4),
                      "difference": round(obs, 4), "p_one_sided_permutation": round(float((np.sum(perm >= obs) + 1) / (N_PERM + 1)), 4),
                      "heavy_rank_among_physical": int(1 + sum(tab[p]["rate"] > tab["HEAVY"]["rate"] for p in PHYSICAL))}
    with open(out, "w") as fh:
        json.dump(res, fh, indent=2, sort_keys=True)
        fh.write("\n")
    print(json.dumps(res, indent=1))


if __name__ == "__main__":
    main()
