"""Recount the accompanying expanded-test-predictions.csv with Python stdlib.

Usage: python recompute-expanded-metrics.py [expanded-test-predictions.csv]
VALID is the positive class. No training, relabeling, or threshold tuning.
AP admits tied scores together; legacy_row_ap preserves CSV row ordering.
"""
import csv
import json
import sys
from collections import defaultdict
from pathlib import Path

path = Path(sys.argv[1]) if len(sys.argv) > 1 else Path(__file__).with_name('expanded-test-predictions.csv')
with path.open(encoding='utf-8-sig', newline='') as stream:
    rows = list(csv.DictReader(stream))
valid = [r for r in rows if r['validity'] == 'VALID']
reject = [r for r in rows if r['validity'] == 'REJECT']
assert len(rows) == len(valid) + len(reject)
def identity_ok(row, model):
    return row['type_name'] == row[f'{model}_type'] and row['identity_id'] == row[f'{model}_identity']

report = {'samples': len(rows), 'valid': len(valid), 'reject': len(reject),
          'identity_correct': {m: sum(identity_ok(r, m) for r in valid) for m in ('m0', 'm1')}}
for field, threshold in [('confidence_score', 1.0), ('learned_validity_score', 0.494140625)]:
    groups = defaultdict(lambda: [0, 0])
    for r in rows:
        g = groups[float(r[field])]
        g[0] += r['validity'] == 'VALID'
        g[1] += 1
    tp = seen = 0
    ap = 0.0
    for score, (pos, count) in sorted(groups.items(), reverse=True):
        tp += pos
        seen += count
        ap += (pos / len(valid)) * (tp / seen)
    positives = [float(r[field]) for r in valid]
    negatives = [float(r[field]) for r in reject]
    auc = sum(1 if p > n else 0.5 if p == n else 0 for p in positives for n in negatives) / (len(valid) * len(reject))
    seen_pos = 0
    legacy_ap = 0.0
    for rank, r in enumerate(sorted(rows, key=lambda r: float(r[field]), reverse=True), 1):
        if r['validity'] == 'VALID':
            seen_pos += 1
            legacy_ap += seen_pos / rank / len(valid)
    accepted = sum(float(r[field]) >= threshold for r in valid)
    rejected = sum(float(r[field]) < threshold for r in reject)
    accepted_correct = sum(float(r[field]) >= threshold and identity_ok(r, 'm1') for r in valid)
    report[field] = {'threshold': threshold, 'valid_accepted': accepted,
                     'false_rejects': len(valid) - accepted, 'reject_rejected': rejected,
                     'false_accepts': len(reject) - rejected,
                     'final_decision_correct': accepted_correct + rejected,
                     'balanced_accuracy': (accepted / len(valid) + rejected / len(reject)) / 2,
                     'auroc': auc, 'grouped_ap': ap, 'legacy_row_ap': legacy_ap}
print(json.dumps(report, indent=2))
