"""Liczy wyniki badania z opublikowanych decyzji (decyzje.jsonl) i zapisuje wyniki.json.

Etap 1 (tylko u autora): jeśli obok leży surowy zapis pelny.jsonl (wszystkie prawdopodobieństwa dla całej
części testowej, około 10 MB, niepublikowany), skrypt najpierw tworzy z niego zwarty decyzje.jsonl dla 1000 par z ids-proba.txt. Zwarty zapis
zachowuje wszystko, czego potrzeba do miar: p_etykiety i suma_p2 są liczone z prawdopodobieństw zaokrąglonych tak
jak w surowym zapisie (4 miejsca).
Etap 2 (każdy): miary z protokołu i analizy po fakcie, wyłącznie z decyzje.jsonl.
"""
import json
import math
import random
import statistics
from collections import Counter, defaultdict
from pathlib import Path

TU = Path(__file__).parent
ZIARNO = 20261002
KURS_NBP = 3.8881  # 2.10.2026, tabela 192/A/NBP/2026
STAWKI = {"clef": 0.24, "clef-flash": 0.09}  # USD za mln tokenów wejścia, Workers AI, odczyt 2.10.2026

ids = [l.strip() for l in open(TU / "ids-proba.txt") if l.strip()]

# Etap 1: surowy zapis -> decyzje.jsonl (pomijany, gdy surowego zapisu nie ma)
if not (TU / "pelny.jsonl").exists():
    print("Brak pelny.jsonl: pomijam etap 1 i liczę wyniki z istniejącego decyzje.jsonl.")
else:
    zbior = set(ids)
    surowe = [w for w in map(json.loads, open(TU / "pelny.jsonl")) if w["id"] in zbior]
    assert len(surowe) == 2 * len(ids) and all(w["status"] == "ok" for w in surowe)
    with open(TU / "decyzje.jsonl", "w") as f:
        for w in sorted(surowe, key=lambda x: (int(x["id"]), x["locale"] != "pl-PL")):
            f.write(json.dumps({
                "id": w["id"], "jezyk": w["locale"], "zdanie": w["utt"], "etykieta": w["gold"],
                "przewidziana": w["pred"], "pewnosc": w["conf"], "p_etykiety": w["probs"].get(w["gold"], 0.0),
                "suma_p2": round(sum(p * p for p in w["probs"].values()), 8),
                "tokeny_wejscia": w["input_tokens"], "sekundy_mac": w["sekundy"],
            }, ensure_ascii=False) + "\n")

# Etap 2: miary z decyzje.jsonl
decyzje = [json.loads(l) for l in open(TU / "decyzje.jsonl")]
po = {l: {d["id"]: d for d in decyzje if d["jezyk"] == l} for l in ("pl-PL", "en-US")}
kolejnosc = sorted(ids, key=int)  # ta sama kolejność co w analiza.py, od niej zależy losowanie bootstrapu
trafna = lambda w: w["przewidziana"] == w["etykieta"]
brier = lambda w: w["suma_p2"] - 2 * w["p_etykiety"] + 1
# 60 intencji MASSIVE: 59 obecnych w części testowej i cooking_query (protokół), bez potrzeby pobierania zbioru
wszystkie_etykiety = sorted({w["etykieta"] for w in decyzje} | {"cooking_query"})


def ece(ws, przedzialy=15):
    kosze = defaultdict(list)
    for w in ws:
        kosze[min(int(w["pewnosc"] * przedzialy), przedzialy - 1)].append(w)
    return sum(len(k) / len(ws) * abs(statistics.mean(trafna(x) for x in k) - statistics.mean(x["pewnosc"] for x in k))
               for k in kosze.values())


def makro_f1(ws, klasy=None):
    tp, fp, fn = Counter(), Counter(), Counter()
    for w in ws:
        if trafna(w):
            tp[w["etykieta"]] += 1
        else:
            fp[w["przewidziana"]] += 1
            fn[w["etykieta"]] += 1
    f1 = []
    for k in (klasy or {w["etykieta"] for w in ws}):
        p = tp[k] / (tp[k] + fp[k]) if tp[k] + fp[k] else 0.0
        r = tp[k] / (tp[k] + fn[k]) if tp[k] + fn[k] else 0.0
        f1.append(2 * p * r / (p + r) if p + r else 0.0)
    return statistics.mean(f1)


def bledy_od_progu(ws, prog=0.9):
    pewne = [w for w in ws if w["pewnosc"] >= prog]
    return sum(not trafna(w) for w in pewne) / len(pewne)


wynik = {"pary": len(ids), "decyzje": len(decyzje), "kurs_nbp": KURS_NBP, "stawki_usd_za_mln": STAWKI, "jezyki": {}}
for l, d in po.items():
    ws = list(d.values())
    tok = [w["tokeny_wejscia"] for w in ws]
    pewne = [w for w in ws if w["pewnosc"] >= 0.9]
    pasmo = [w for w in ws if 0.9 <= w["pewnosc"] < 0.95]
    wynik["jezyki"][l] = {
        "trafnosc_intencji": round(statistics.mean(trafna(w) for w in ws), 4),
        "trafnosc_obszaru": round(statistics.mean(w["przewidziana"].split("_")[0] == w["etykieta"].split("_")[0] for w in ws), 4),
        "makro_f1": round(makro_f1(ws), 4),
        "ece": round(ece(ws), 4),
        "brier": round(statistics.mean(brier(w) for w in ws), 4),
        "srednia_pewnosc": round(statistics.mean(w["pewnosc"] for w in ws), 4),
        "tokeny_srednio": round(statistics.mean(tok), 2), "tokeny_mediana": statistics.median(tok),
        "sekundy_mac_mediana": round(statistics.median(w["sekundy_mac"] for w in ws), 2),
        "koszt_10_tys_decyzji_zl": {m: round(statistics.mean(tok) * 10_000 * s / 1e6 * KURS_NBP, 2) for m, s in STAWKI.items()},
        "po_fakcie": {
            "pewnosc_od_0_9": {"decyzje": len(pewne), "bledne": sum(not trafna(w) for w in pewne),
                               "trafnosc": round(statistics.mean(trafna(w) for w in pewne), 4)},
            "ece_10_przedzialow": round(ece(ws, 10), 4),
            "makro_f1_60_etykiet": round(makro_f1(ws, wszystkie_etykiety), 4),
            "pewnosc_0_90_do_0_95": {"decyzje": len(pasmo), "srednia_pewnosc": round(statistics.mean(w["pewnosc"] for w in pasmo), 4),
                                     "trafnosc": round(statistics.mean(trafna(w) for w in pasmo), 4)},
        },
    }

pl = [trafna(po["pl-PL"][i]) for i in kolejnosc]
en = [trafna(po["en-US"][i]) for i in kolejnosc]
r = random.Random(ZIARNO)
n = len(kolejnosc)
roznice = sorted(sum(pl[j] - en[j] for j in (r.randrange(n) for _ in range(n))) / n for _ in range(10_000))
b = sum(1 for p, e in zip(pl, en) if p and not e)
c = sum(1 for p, e in zip(pl, en) if e and not p)
p_mc = min(1.0, 2 * sum(math.comb(b + c, i) for i in range(min(b, c) + 1)) / 2 ** (b + c))
wynik["roznica_pl_en"] = {
    "punkty_proc": round(100 * (statistics.mean(pl) - statistics.mean(en)), 1),
    "ci95_punkty_proc": [round(100 * roznice[249], 1), round(100 * roznice[9749], 1)],
    "tylko_pl_trafne": b, "tylko_en_trafne": c, "mcnemar_p": round(p_mc, 3),
    "ta_sama_decyzja": sum(po["pl-PL"][i]["przewidziana"] == po["en-US"][i]["przewidziana"] for i in kolejnosc),
}

# Analizy po fakcie (poza protokołem): osobne losowanie z tym samym ziarnem
r = random.Random(ZIARNO)
rb, rp = [], []
for _ in range(10_000):
    s = [kolejnosc[r.randrange(n)] for _ in range(n)]
    rb.append(statistics.mean(brier(po["pl-PL"][i]) - brier(po["en-US"][i]) for i in s))
    rp.append(bledy_od_progu([po["pl-PL"][i] for i in s]) - bledy_od_progu([po["en-US"][i] for i in s]))
rb.sort()
rp.sort()
r = random.Random(ZIARNO)
rf = []
for _ in range(2_000):
    s = [kolejnosc[r.randrange(n)] for _ in range(n)]
    rf.append(makro_f1([po["pl-PL"][i] for i in s]) - makro_f1([po["en-US"][i] for i in s]))
rf.sort()
licz = Counter(po["pl-PL"][i]["etykieta"] for i in kolejnosc)
rzadkie = [i for i in kolejnosc if licz[po["pl-PL"][i]["etykieta"]] <= 11]
czeste = [i for i in kolejnosc if licz[po["pl-PL"][i]["etykieta"]] > 11]
wynik["po_fakcie"] = {
    "brier_roznica_pl_en": round(statistics.mean(brier(po["pl-PL"][i]) - brier(po["en-US"][i]) for i in kolejnosc), 4),
    "brier_roznica_ci95": [round(rb[249], 3), round(rb[9749], 3)],
    "bledy_od_0_9_roznica_punkty_proc": round(100 * (bledy_od_progu(list(po["pl-PL"].values())) - bledy_od_progu(list(po["en-US"].values()))), 1),
    "bledy_od_0_9_ci95_punkty_proc": [round(100 * rp[249], 1), round(100 * rp[9749], 1)],
    "makro_f1_roznica_pl_en": round(makro_f1(list(po["pl-PL"].values())) - makro_f1(list(po["en-US"].values())), 4),
    "makro_f1_roznica_ci95": [round(rf[49], 3), round(rf[1949], 3)],
    "brier_bledy_i_pewnosc": {l: {"bledne": sum(not trafna(w) for w in po[l].values()),
                                  "pewnosc_przy_bledach": round(statistics.mean(w["pewnosc"] for w in po[l].values() if not trafna(w)), 3),
                                  "brier_z_blednych": round(sum(brier(w) for w in po[l].values() if not trafna(w)) / n, 4),
                                  "brier_z_trafnych": round(sum(brier(w) for w in po[l].values() if trafna(w)) / n, 4)}
                              for l in ("pl-PL", "en-US")},
    "intencje_do_11_przykladow": {"pary": len(rzadkie), "trafnosc_pl": round(statistics.mean(trafna(po["pl-PL"][i]) for i in rzadkie), 4),
                                  "trafnosc_en": round(statistics.mean(trafna(po["en-US"][i]) for i in rzadkie), 4),
                                  "tylko_pl_trafne": sum(trafna(po["pl-PL"][i]) and not trafna(po["en-US"][i]) for i in rzadkie),
                                  "tylko_en_trafne": sum(trafna(po["en-US"][i]) and not trafna(po["pl-PL"][i]) for i in rzadkie)},
    "progi_pewnosci": {str(prog): {l: {"decyzje": sum(w["pewnosc"] >= prog for w in po[l].values()),
                                         "bledne": sum(w["pewnosc"] >= prog and not trafna(w) for w in po[l].values())}
                                     for l in ("pl-PL", "en-US")} for prog in (0.8, 0.85, 0.9, 0.95, 0.97)},
    "intencje_powyzej_11": {"pary": len(czeste), "trafnosc_pl": round(statistics.mean(trafna(po["pl-PL"][i]) for i in czeste), 4),
                            "trafnosc_en": round(statistics.mean(trafna(po["en-US"][i]) for i in czeste), 4)},
}
wynik["schematy"] = json.load(open(TU / "tokeny-schematu.json"))["jezyki"]
json.dump(wynik, open(TU / "wyniki.json", "w"), ensure_ascii=False, indent=2)
print(json.dumps({k: wynik[k] for k in ("roznica_pl_en", "po_fakcie")}, ensure_ascii=False, indent=1))
print(json.dumps({l: v["po_fakcie"] for l, v in wynik["jezyki"].items()}, ensure_ascii=False, indent=1))
