#!/usr/bin/env python3
"""Preregistered NYC wastewater lag7 vs lag0 rise-prediction AUROC."""
from __future__ import annotations

import csv
import json
import math
from collections import defaultdict
from datetime import date, datetime, timedelta
from pathlib import Path

ROOT = Path(__file__).resolve().parents[1]
DATA = ROOT / "data"
OUT = ROOT / "results"

WINDOW_START = date(2023, 4, 1)
WINDOW_END = date(2025, 6, 30)
LAG = 7
HORIZON = 7
RISE_RATIO = 1.10
MIN_CASES = 50.0


def parse_ymd(s: str) -> date:
    return date.fromisoformat(s[:10])


def parse_mdy(s: str) -> date:
    return datetime.strptime(s, "%m/%d/%Y").date()


def auroc(scores: list[float], labels: list[int]) -> float:
    """Mann–Whitney AUROC with midranks for ties. Deterministic."""
    n_pos = sum(1 for y in labels if y == 1)
    n_neg = sum(1 for y in labels if y == 0)
    if n_pos == 0 or n_neg == 0:
        return float("nan")
    # rank scores ascending with midranks
    order = sorted(range(len(scores)), key=lambda i: scores[i])
    ranks = [0.0] * len(scores)
    i = 0
    n = len(scores)
    while i < n:
        j = i
        while j + 1 < n and scores[order[j + 1]] == scores[order[i]]:
            j += 1
        # ranks are 1-based; midrank average of i+1 .. j+1
        mid = 0.5 * ((i + 1) + (j + 1))
        for k in range(i, j + 1):
            ranks[order[k]] = mid
        i = j + 1
    sum_ranks_pos = sum(ranks[i] for i, y in enumerate(labels) if y == 1)
    # U = sum_ranks_pos - n_pos*(n_pos+1)/2
    u = sum_ranks_pos - n_pos * (n_pos + 1) / 2.0
    return u / (n_pos * n_neg)


def load_ww() -> dict[date, float]:
    by_day: dict[date, list[tuple[float, float]]] = defaultdict(list)
    with open(DATA / "nwss_nyc_percentile.csv", newline="") as f:
        for r in csv.DictReader(f):
            d = parse_ymd(r["date_end"])
            pop = float(r["population_served"])
            perc = float(r["percentile"])
            by_day[d].append((pop, perc))
    out: dict[date, float] = {}
    for d, pairs in by_day.items():
        wsum = sum(p for p, _ in pairs)
        if wsum <= 0:
            continue
        out[d] = sum(p * v for p, v in pairs) / wsum
    return out


def load_cases() -> dict[date, float]:
    out: dict[date, float] = {}
    with open(DATA / "nyc_cases_7day.csv", newline="") as f:
        for r in csv.DictReader(f):
            d = parse_mdy(r["date_of_interest"])
            out[d] = float(r["CASE_COUNT_7DAY_AVG"] or 0.0)
    return out


def main() -> None:
    ww = load_ww()
    cases = load_cases()
    scores0: list[float] = []
    scores7: list[float] = []
    labels: list[int] = []
    rows_out = []

    t = WINDOW_START
    while t <= WINDOW_END:
        t7 = t - timedelta(days=LAG)
        th = t + timedelta(days=HORIZON)
        if t not in ww or t7 not in ww:
            t += timedelta(days=1)
            continue
        if t not in cases or th not in cases:
            t += timedelta(days=1)
            continue
        c0 = cases[t]
        if c0 < MIN_CASES:
            t += timedelta(days=1)
            continue
        rise = 1 if (cases[th] / c0) >= RISE_RATIO else 0
        s0 = ww[t]
        s7 = ww[t7]
        scores0.append(s0)
        scores7.append(s7)
        labels.append(rise)
        rows_out.append(
            {
                "date": t.isoformat(),
                "ww_lag0": s0,
                "ww_lag7": s7,
                "cases": c0,
                "cases_horizon": cases[th],
                "rise": rise,
            }
        )
        t += timedelta(days=1)

    a0 = auroc(scores0, labels)
    a7 = auroc(scores7, labels)
    n_pos = sum(labels)
    n_neg = len(labels) - n_pos
    delta = a7 - a0
    success = bool(delta > 0 and not math.isnan(delta))

    OUT.mkdir(parents=True, exist_ok=True)
    result = {
        "auroc_lag0": round(a0, 6),
        "auroc_lag7": round(a7, 6),
        "delta_auroc": round(delta, 6),
        "n_days": len(labels),
        "n_pos": n_pos,
        "n_neg": n_neg,
        "lag_days": LAG,
        "horizon_days": HORIZON,
        "rise_ratio": RISE_RATIO,
        "min_cases": MIN_CASES,
        "window_start": WINDOW_START.isoformat(),
        "window_end": WINDOW_END.isoformat(),
        "success_lag_beats_contemporaneous": success,
    }
    (OUT / "R1.json").write_text(json.dumps(result, indent=2, sort_keys=True) + "\n")
    with open(OUT / "daily.csv", "w", newline="") as f:
        w = csv.DictWriter(
            f,
            fieldnames=["date", "ww_lag0", "ww_lag7", "cases", "cases_horizon", "rise"],
        )
        w.writeheader()
        for r in rows_out:
            w.writerow(r)
    print(json.dumps(result, indent=2, sort_keys=True))


if __name__ == "__main__":
    main()
