"""Analiza wyników pomiaru Clef-flash według protokol.md.

Użycie: python analiza.py pelny.jsonl [powtorka.jsonl]
"""
import json
import math
import random
import statistics
import sys
from collections import Counter, defaultdict

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


def wczytaj(sciezka):
    return [json.loads(l) for l in open(sciezka)]


def makro_f1(pary):
    klasy = {g for g, _ in pary}
    tp, fp, fn = Counter(), Counter(), Counter()
    for g, p in pary:
        if g == p:
            tp[g] += 1
        else:
            fp[p] += 1
            fn[g] += 1
    wyniki = []
    for k in klasy:
        prec = tp[k] / (tp[k] + fp[k]) if tp[k] + fp[k] else 0.0
        rec = tp[k] / (tp[k] + fn[k]) if tp[k] + fn[k] else 0.0
        wyniki.append(2 * prec * rec / (prec + rec) if prec + rec else 0.0)
    return sum(wyniki) / len(wyniki)


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


def brier(wiersze):
    return statistics.mean(sum((p - (o == w["gold"])) ** 2 for o, p in w["probs"].items()) for w in wiersze)


def mcnemar(b, c):
    n = b + c
    if n == 0:
        return 1.0
    k = min(b, c)
    return min(1.0, 2 * sum(math.comb(n, i) for i in range(k + 1)) / 2 ** n)


def main():
    wiersze = wczytaj(sys.argv[1])
    ok = [w for w in wiersze if w["status"] == "ok"]
    print(f"rekordy: {len(wiersze)}, ok: {len(ok)}, nie zmierzono (błąd): {len(wiersze) - len(ok)}")
    po_jezyku = {l: {w["id"]: w for w in ok if w["locale"] == l} for l in ("pl-PL", "en-US")}
    wspolne = sorted(set(po_jezyku["pl-PL"]) & set(po_jezyku["en-US"]), key=int)
    print(f"pary z wynikiem w obu językach: {len(wspolne)}")

    for l, d in po_jezyku.items():
        w = [d[i] for i in wspolne]
        traf = statistics.mean(x["pred"] == x["gold"] for x in w)
        scen = statistics.mean(x["pred"].split("_")[0] == x["gold"].split("_")[0] for x in w)
        tok = [x["input_tokens"] for x in w]
        print(f"\n{l}: trafność intencji {traf:.4f}, scenariusza {scen:.4f}, makro-F1 {makro_f1([(x['gold'], x['pred']) for x in w]):.4f}")
        print(f"  ECE {ece(w):.4f}, Brier {brier(w):.4f}, średnia pewność {statistics.mean(x['conf'] for x in w):.4f}")
        print(f"  tokeny wejścia: średnia {statistics.mean(tok):.1f}, mediana {statistics.median(tok)}, min {min(tok)}, max {max(tok)}")
        print(f"  czas na decyzję (Mac M5 Pro, MPS): mediana {statistics.median(x['sekundy'] for x in w):.2f} s")
        for model, stawka in STAWKI.items():
            usd = statistics.mean(tok) * 10_000 * stawka / 1e6
            print(f"  koszt 10 tys. decyzji, {model}: {usd:.4f} USD = {usd * KURS_NBP:.4f} zł")

    pl = [po_jezyku["pl-PL"][i]["pred"] == po_jezyku["pl-PL"][i]["gold"] for i in wspolne]
    en = [po_jezyku["en-US"][i]["pred"] == po_jezyku["en-US"][i]["gold"] for i in wspolne]
    r = random.Random(ZIARNO)
    n = len(wspolne)
    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)
    print(f"\nPL−EN: {statistics.mean(pl) - statistics.mean(en):+.4f}, 95% CI [{roznice[249]:+.4f}, {roznice[9749]:+.4f}]")
    print(f"McNemar: tylko PL trafne {b}, tylko EN trafne {c}, p = {mcnemar(b, c):.4g}")
    zgodne = sum(po_jezyku['pl-PL'][i]['pred'] == po_jezyku['en-US'][i]['pred'] for i in wspolne)
    print(f"ta sama decyzja w obu językach: {zgodne}/{n}")
    tok_pl = statistics.mean(po_jezyku["pl-PL"][i]["input_tokens"] for i in wspolne)
    tok_en = statistics.mean(po_jezyku["en-US"][i]["input_tokens"] for i in wspolne)
    print(f"tokeny PL/EN (całe zapytanie): {tok_pl / tok_en:.3f}")

    if len(sys.argv) > 2:
        drugi = {(w["id"], w["locale"]): w for w in wczytaj(sys.argv[2]) if w["status"] == "ok"}
        wspolne2 = [w for w in ok if (w["id"], w["locale"]) in drugi]
        zgodne = sum(1 for w in wspolne2 if drugi[(w["id"], w["locale"])]["pred"] == w["pred"])
        maks = max(abs(drugi[(w["id"], w["locale"])]["probs"][o] - p) for w in wspolne2 for o, p in w["probs"].items())
        print(f"\nporównanie z drugim plikiem: ta sama decyzja top-1 w {zgodne}/{len(wspolne2)}, maks. różnica prawdopodobieństw {maks:.6f}")


if __name__ == "__main__":
    main()
