"""Synthetic precision/recall tutorial; no trained model or biological results."""
import argparse

import numpy as np
from sklearn.metrics import average_precision_score, precision_recall_curve


def checked_arrays(labels, scores):
    y = np.asarray(labels)
    s = np.asarray(scores, dtype=float)
    if y.ndim != 1 or s.ndim != 1 or y.size != s.size or not y.size:
        raise ValueError("Labels and scores must be nonempty 1D arrays of equal length")
    if not np.isin(y, [0, 1]).all() or np.unique(y).size != 2:
        raise ValueError("Labels must contain both 0 and 1")
    if not np.isfinite(s).all():
        raise ValueError("Scores must be finite")
    return y.astype(int), s


def choose_threshold(labels, scores, min_precision=0.60):
    y, s = checked_arrays(labels, scores)
    if not np.isfinite(min_precision) or not 0 < min_precision <= 1:
        raise ValueError("min_precision must be in (0, 1]")
    precision, recall, thresholds = precision_recall_curve(y, s, pos_label=1)
    # The final plotting endpoint has no threshold and must not be selected.
    eligible = np.flatnonzero(precision[:-1] >= min_precision)
    if not eligible.size:
        raise ValueError("No observed threshold meets the precision requirement")
    # Maximize recall, then precision, then threshold for a deterministic tie rule.
    best = max(eligible, key=lambda i: (recall[i], precision[i], thresholds[i]))
    return float(thresholds[best])


def report(name, labels, scores, threshold):
    y, s = checked_arrays(labels, scores)
    if not np.isfinite(threshold):
        raise ValueError("Threshold must be finite")
    pred = s >= threshold
    tp = int(np.sum(pred & (y == 1)))
    fp = int(np.sum(pred & (y == 0)))
    fn = int(np.sum(~pred & (y == 1)))
    tn = int(np.sum(~pred & (y == 0)))
    precision = f"{tp / (tp + fp):.3f}" if tp + fp else "undefined"
    print(f"{name}: threshold={threshold:.3f} selected={tp + fp}/{len(y)}")
    print(f"  TP={tp} FP={fp} FN={fn} TN={tn}")
    print(f"  precision={precision} recall={tp / (tp + fn):.3f}")
    print(f"  AP={average_precision_score(y, s):.3f} prevalence={y.mean():.3f}")


def demo_data():
    # Separate invented validation and test records, ordered by decreasing score.
    scores = np.r_[
        [.96, .91, .87, .82, .78, .73, .68, .62, .57, .52,
         .47, .42, .37, .32, .28, .24, .20, .16, .12, .08],
        np.linspace(.07, .001, 80),
    ]
    valid = np.zeros(100, dtype=int)
    test = np.zeros(100, dtype=int)
    valid[np.array([1, 3, 4, 6, 9, 11, 16, 20]) - 1] = 1
    test[np.array([1, 2, 5, 8, 12, 20, 24, 30]) - 1] = 1
    return valid, scores.copy(), test, scores.copy()


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--min-precision", type=float, default=0.60)
    args = parser.parse_args()
    valid_y, valid_s, test_y, test_s = demo_data()
    try:
        threshold = choose_threshold(valid_y, valid_s, args.min_precision)
    except ValueError as exc:
        parser.error(str(exc))
    print(f"Validation precision requirement: {args.min_precision:.3f}")
    report("validation", valid_y, valid_s, threshold)
    report("test (frozen threshold)", test_y, test_s, threshold)
    report("test (predeclared 0.5 reference)", test_y, test_s, 0.5)
    print(f"Always-negative test accuracy: {(test_y == 0).mean():.3f}")


if __name__ == "__main__":
    main()
