"""Synthetic candidate selection; no model, network, or paid calls."""
import argparse
import csv
import gzip
import hashlib
import json
import math
from pathlib import Path
import platform
import random
import sys
import time

N_VALUES = (1, 2, 4, 8, 16)
GEN_COST, VERIFY_COST = 10, 2
# name, marginal candidate correctness, sensitivity, false-positive rate, dead fraction
SCENARIOS = (
    ("independent", 0.4, 0.9, 0.1, 0.0),
    ("weaker_verifier", 0.4, 0.9, 0.35, 0.0),
    ("shared_failure", 0.4, 0.9, 0.1, 0.5),
    ("perfect_verifier", 0.4, 1.0, 0.0, 0.0),
)


def theory(p, a, b, dead, n):
    assert 0 <= dead < 1 and 0 <= p <= 1 - dead
    assert 0 <= a <= 1 and 0 <= b <= 1 and n >= 1
    ans = dict(success=0.0, wrong=0.0, abstain=0.0,
               oracle=0.0, attempts=0.0)
    for weight, conditional_p in ((dead, 0.0), (1-dead, p/(1-dead))):
        q = conditional_p*a + (1-conditional_p)*b
        r = 1-q
        series = sum(r**k for k in range(n))
        ans["success"] += weight*conditional_p*a*series
        ans["wrong"] += weight*(1-conditional_p)*b*series
        ans["abstain"] += weight*r**n
        ans["oracle"] += weight*(1-(1-conditional_p)**n)
        ans["attempts"] += weight*series
    ans["coverage"] = 1-ans["abstain"]
    ans["accepted_accuracy"] = (ans["success"]/ans["coverage"]
                                if ans["coverage"] else None)
    ans["early_cost"] = (GEN_COST+VERIFY_COST)*ans["attempts"]
    ans["batch_cost"] = (GEN_COST+VERIFY_COST)*n
    assert abs(ans["success"]+ans["wrong"]+ans["abstain"]-1) < 1e-12
    return ans


def wilson(k, total):
    if not total:
        return None, None
    z = 1.959963984540054
    rate = k/total
    den = 1+z*z/total
    center = (rate+z*z/(2*total))/den
    half = z*math.sqrt(rate*(1-rate)/total+z*z/(4*total*total))/den
    return max(0.0, center-half), min(1.0, center+half)


def mean_se(total, total_sq, count):
    mean = total/count
    variance = max(0.0, (total_sq-total*total/count)/(count-1))
    return mean, math.sqrt(variance/count)


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--seed", type=int, default=20261007)
    parser.add_argument("--trials", type=int, default=5000)
    parser.add_argument("--out", required=True)
    args = parser.parse_args()
    if not 100 <= args.trials <= 20000:
        parser.error("Use 100 to 20000 independent request trials.")
    out = Path(args.out)
    if any(p.is_symlink() for p in (out, *out.parents)):
        raise ValueError('Output path must not traverse symlinks')
    out.mkdir(parents=True, exist_ok=False)  # Preserve earlier runs.
    start = time.perf_counter()
    rng = random.Random(args.seed)
    names = ("success", "wrong", "abstain", "oracle", "baseline",
             "attempts", "attempts_sq", "delta", "delta_sq")
    totals = {(s[0], n): dict.fromkeys(names, 0) for s in SCENARIOS
              for n in N_VALUES}
    raw_path = out / "all_trials.jsonl.gz"
    with gzip.open(raw_path, "wt", encoding="utf-8") as raw:
        for trial in range(args.trials):
            # Common random numbers pair scenarios and sample caps.
            latent_u = rng.random()
            truth_u = [rng.random() for _ in range(max(N_VALUES))]
            verify_u = [rng.random() for _ in range(max(N_VALUES))]
            for name, p, a, b, dead in SCENARIOS:
                blocked = latent_u < dead
                p_good = p/(1-dead)
                truth = [int(not blocked and u < p_good) for u in truth_u]
                accepted = [int(u < (a if c else b))
                            for u, c in zip(verify_u, truth)]
                raw.write(json.dumps(dict(
                    trial=trial, scenario=name, blocked=blocked,
                    correct=truth, accepted=accepted)) + "\n")
                first_success = int(accepted[0] and truth[0])
                for n in N_VALUES:
                    # Selection sees verifier decisions, never truth labels.
                    chosen = next((i for i in range(n) if accepted[i]), None)
                    success = int(chosen is not None and truth[chosen] == 1)
                    wrong = int(chosen is not None and truth[chosen] == 0)
                    abstain = int(chosen is None)
                    attempts = n if chosen is None else chosen+1
                    delta = success-first_success
                    metrics = (success, wrong, abstain, int(any(truth[:n])),
                               truth[0], attempts, attempts**2, delta, delta**2)
                    acc = totals[name, n]
                    for key, value in zip(names, metrics):
                        acc[key] += value
    rows = []
    for name, p, a, b, dead in SCENARIOS:
        for n in N_VALUES:
            t = totals[name, n]
            count = args.trials
            assert t["success"]+t["wrong"]+t["abstain"] == count
            assert t["success"] <= t["oracle"]
            if name == "perfect_verifier":
                assert t["success"] == t["oracle"] and t["wrong"] == 0
            expected = theory(p, a, b, dead, n)
            accepted_count = count-t["abstain"]
            lo, hi = wilson(t["success"], count)
            alo, ahi = wilson(t["success"], accepted_count)
            attempts, attempts_se = mean_se(t["attempts"], t["attempts_sq"], count)
            delta, delta_se = mean_se(t["delta"], t["delta_sq"], count)
            row = dict(scenario=name, n=n, trials=count,
                       success=t["success"]/count, success_lo=lo, success_hi=hi,
                       wrong=t["wrong"]/count, abstain=t["abstain"]/count,
                       coverage=accepted_count/count,
                       accepted_accuracy=(t["success"]/accepted_count
                                          if accepted_count else None),
                       accepted_accuracy_lo=alo, accepted_accuracy_hi=ahi,
                       oracle=t["oracle"]/count,
                       unverified_baseline=t["baseline"]/count,
                       early_cost=(GEN_COST+VERIFY_COST)*attempts,
                       early_cost_se=(GEN_COST+VERIFY_COST)*attempts_se,
                       batch_cost=(GEN_COST+VERIFY_COST)*n,
                       delta_vs_n1=delta, delta_se=delta_se)
            row.update({"expected_"+k: v for k, v in expected.items()})
            rows.append(row)
    with (out / "summary.csv").open("w", newline="", encoding="utf-8") as f:
        writer = csv.DictWriter(f, fieldnames=list(rows[0]))
        writer.writeheader()
        writer.writerows(rows)
    manifest = dict(seed=args.seed, trials=args.trials, sample_caps=N_VALUES,
                    scenarios=SCENARIOS, generation_units=GEN_COST,
                    verification_units=VERIFY_COST, python=sys.version,
                    platform=platform.platform(),
                    source_sha256=hashlib.sha256(Path(__file__).read_bytes()).hexdigest(),
                    raw_sha256=hashlib.sha256(raw_path.read_bytes()).hexdigest(),
                    simulation_elapsed_seconds=time.perf_counter()-start,
                    cost_units="invented dimensionless policy units",
                    real_model_experiment="not run")
    (out / "manifest.json").write_text(
        json.dumps(manifest, indent=2), encoding="utf-8")
    for row in rows:
        print(row["scenario"], row["n"],
              "success", round(row["success"], 4),
              "expected", round(row["expected_success"], 4),
              "early units", round(row["early_cost"], 2))


if __name__ == "__main__":
    assert abs(theory(.4, .9, .1, 0, 4)["success"]-.76014432) < 1e-12
    assert abs(theory(.4, .9, .1, 0, 4)["early_cost"]-25.338144) < 1e-12
    assert abs(theory(.4, .9, .1, .5, 4)["oracle"]-.4992) < 1e-12
    assert theory(.4, 0, 0, 0, 4)["abstain"] == 1
    main()
