#!/usr/bin/env python3
"""Simulation-based recoverability check for the Genesis pedigree test."""

from __future__ import annotations

import os
from concurrent.futures import ProcessPoolExecutor, as_completed
from pathlib import Path

import numpy as np
import pandas as pd

from run_analysis import N_GRID, fit_training


ROOT = Path(__file__).resolve().parent
RESULTS = ROOT / "results"
AMBIENT = 80
RETENTION = 0.90
TRUE_NS = [5, 20, 100]
TRUE_SIGMAS = [10, 30, 50]
REPLICATES = 100


def category(n: int) -> str:
    if n <= 5:
        return "few"
    if n <= 50:
        return "tens"
    return "many"


def one_replicate(args):
    true_n, true_sigma, rep, seed = args
    rng = np.random.default_rng(seed)
    d = 600 - AMBIENT
    y = [600.0]
    k = true_n
    for _ in range(7):
        k = rng.binomial(k, RETENTION)
        y.append(AMBIENT + d * k / true_n + rng.normal(0, true_sigma))
    # Held-out placeholders are unused by fit_training.
    y.extend([np.nan, np.nan])
    y = np.asarray(y, dtype=float)
    fits = []
    for candidate_n in N_GRID:
        z = fit_training(y, AMBIENT, candidate_n)
        fits.append((candidate_n, z["train_loglik"], z["retention"], z["sigma"]))
    best = max(fits, key=lambda x: x[1])
    maxll = best[1]
    supported = [n for n, ll, _, _ in fits if maxll - ll <= 1.92]
    supported_categories = sorted(set(category(n) for n in supported))
    return {
        "true_N": true_n,
        "true_category": category(true_n),
        "true_sigma": true_sigma,
        "replicate": rep,
        "best_N": best[0],
        "best_category": category(best[0]),
        "best_retention": best[2],
        "best_sigma": best[3],
        "category_correct": category(best[0]) == category(true_n),
        "true_category_in_profile_support": category(true_n) in supported_categories,
        "supported_categories": ";".join(supported_categories),
        "training_values": ";".join(f"{v:.3f}" for v in y[:8]),
    }


def main():
    jobs = []
    ss = np.random.SeedSequence(20260809)
    child_seeds = ss.spawn(len(TRUE_NS) * len(TRUE_SIGMAS) * REPLICATES)
    idx = 0
    for true_n in TRUE_NS:
        for true_sigma in TRUE_SIGMAS:
            for rep in range(REPLICATES):
                jobs.append((true_n, true_sigma, rep, int(child_seeds[idx].generate_state(1)[0])))
                idx += 1

    rows = []
    workers = min(8, max(1, (os.cpu_count() or 2) - 1))
    with ProcessPoolExecutor(max_workers=workers) as pool:
        futures = [pool.submit(one_replicate, job) for job in jobs]
        for i, future in enumerate(as_completed(futures), 1):
            rows.append(future.result())
            if i % 100 == 0:
                print(f"completed {i}/{len(jobs)}", flush=True)

    raw = pd.DataFrame(rows).sort_values(["true_N", "true_sigma", "replicate"])
    raw.to_csv(RESULTS / "parameter_recovery_raw.csv", index=False)

    summary = (
        raw.groupby(["true_N", "true_category", "true_sigma"])
        .agg(
            replicates=("replicate", "count"),
            exact_category_recovery=("category_correct", "mean"),
            true_category_in_95pct_profile=("true_category_in_profile_support", "mean"),
            median_best_N=("best_N", "median"),
            median_best_retention=("best_retention", "median"),
            median_best_sigma=("best_sigma", "median"),
        )
        .reset_index()
    )
    summary.to_csv(RESULTS / "parameter_recovery_summary.csv", index=False)

    confusion = pd.crosstab(
        [raw.true_category, raw.true_sigma], raw.best_category, normalize="index"
    ).reset_index()
    confusion.to_csv(RESULTS / "parameter_recovery_confusion.csv", index=False)
    print(summary.to_string(index=False))
    print("\nConfusion proportions:\n", confusion.to_string(index=False))


if __name__ == "__main__":
    main()
